191 lines
5.3 KiB
Python
191 lines
5.3 KiB
Python
from __future__ import annotations
|
|
|
|
from copy import deepcopy
|
|
from typing import Annotated
|
|
|
|
import pytest
|
|
from pydantic import BaseModel, Field
|
|
|
|
from tests.authoring.helpers import (
|
|
AppendState,
|
|
DefaultedState,
|
|
NestedWorkflowState,
|
|
TypedDictInput,
|
|
WorkflowInput,
|
|
WorkflowOutput,
|
|
WorkflowState,
|
|
)
|
|
from wf_authoring import WorkflowBuilder, state_field
|
|
from wf_core.paths import StatePath
|
|
|
|
|
|
class DotAliasState(BaseModel):
|
|
person_tags: Annotated[
|
|
list[str],
|
|
Field(alias="person.name"),
|
|
state_field(reducer="wf.std.append"),
|
|
]
|
|
|
|
|
|
class RevisedInput(BaseModel):
|
|
query: str
|
|
|
|
|
|
class RevisedState(BaseModel):
|
|
result: str | None = None
|
|
|
|
|
|
class RevisedOutput(BaseModel):
|
|
result: str
|
|
|
|
|
|
def test_builder_accepts_basemodel_classes_for_workflow_schemas() -> None:
|
|
builder = WorkflowBuilder(
|
|
name="model_schema_demo",
|
|
input_schema=WorkflowInput,
|
|
state_schema=WorkflowState,
|
|
output_schema=WorkflowOutput,
|
|
start="start",
|
|
)
|
|
|
|
workflow = builder.compile()
|
|
|
|
assert workflow.input_schema.properties["text"]["type"] == "string"
|
|
assert workflow.output_schema.properties["text"]["type"] == "string"
|
|
fields = workflow.state_schema.field_map()
|
|
assert set(fields) == {"text", "count", "tags"}
|
|
assert fields["text"].type == "string"
|
|
assert fields["count"].type == "integer"
|
|
assert fields["tags"].type == "array"
|
|
|
|
|
|
def test_set_contract_replaces_selected_fields_from_python_models() -> None:
|
|
builder = WorkflowBuilder(
|
|
name="editable_contract",
|
|
input_schema=WorkflowInput,
|
|
state_schema=WorkflowState,
|
|
output_schema=WorkflowOutput,
|
|
outcomes=("ok",),
|
|
start="start",
|
|
)
|
|
original_input = deepcopy(builder.input_schema)
|
|
|
|
builder.set_contract(
|
|
state_schema=RevisedState,
|
|
output_schema=RevisedOutput,
|
|
outcomes=("completed", "rejected"),
|
|
)
|
|
workflow = builder.compile()
|
|
|
|
assert workflow.input_schema == original_input
|
|
assert set(workflow.state_schema.field_map()) == {"result"}
|
|
assert workflow.output_schema.properties["result"]["type"] == "string"
|
|
assert workflow.outcomes == ["completed", "rejected"]
|
|
|
|
|
|
def test_set_contract_replacement_is_atomic_when_normalization_fails() -> None:
|
|
builder = WorkflowBuilder(
|
|
name="atomic_contract",
|
|
input_schema=WorkflowInput,
|
|
state_schema=WorkflowState,
|
|
output_schema=WorkflowOutput,
|
|
outcomes=("ok",),
|
|
start="start",
|
|
)
|
|
original = builder.compile()
|
|
|
|
with pytest.raises(ValueError):
|
|
builder.set_contract(
|
|
input_schema=RevisedInput,
|
|
output_schema={"type": "definitely-not-a-json-schema-type"},
|
|
)
|
|
|
|
assert builder.compile() == original
|
|
|
|
|
|
def test_builder_accepts_typeddict_for_json_schema_refs() -> None:
|
|
builder = WorkflowBuilder(
|
|
name="typed_dict_schema_demo",
|
|
input_schema=TypedDictInput,
|
|
state_schema=WorkflowState,
|
|
output_schema=WorkflowOutput,
|
|
start="start",
|
|
)
|
|
|
|
workflow = builder.compile()
|
|
|
|
assert workflow.input_schema.properties["text"]["type"] == "string"
|
|
|
|
|
|
def test_state_basemodel_can_declare_reducer_with_annotated_metadata() -> None:
|
|
builder = WorkflowBuilder(
|
|
name="state_metadata_demo",
|
|
input_schema=WorkflowInput,
|
|
state_schema=AppendState,
|
|
output_schema=WorkflowOutput,
|
|
start="start",
|
|
)
|
|
|
|
workflow = builder.compile()
|
|
|
|
fields = workflow.state_schema.field_map()
|
|
assert fields["items"].type == "array"
|
|
assert fields["items"].reducer.name == "wf.std.append"
|
|
|
|
|
|
def test_state_basemodel_seeds_safe_initial_defaults() -> None:
|
|
builder = WorkflowBuilder(
|
|
name="state_defaults_demo",
|
|
input_schema=WorkflowInput,
|
|
state_schema=DefaultedState,
|
|
output_schema=WorkflowOutput,
|
|
start="start",
|
|
)
|
|
|
|
workflow = builder.compile()
|
|
|
|
fields = workflow.state_schema.field_map()
|
|
assert fields["items"].default == []
|
|
assert fields["metadata"].default == {}
|
|
assert fields["explicit"].default == 3
|
|
|
|
|
|
def test_nested_state_basemodel_projects_parent_and_child_paths() -> None:
|
|
builder = WorkflowBuilder(
|
|
name="nested_state_schema_demo",
|
|
input_schema=WorkflowInput,
|
|
state_schema=NestedWorkflowState,
|
|
output_schema=WorkflowOutput,
|
|
start="start",
|
|
)
|
|
|
|
workflow = builder.compile()
|
|
|
|
fields = workflow.state_schema.field_map()
|
|
assert set(fields) == {
|
|
"person",
|
|
"person.name",
|
|
"person.tags",
|
|
}
|
|
assert fields["person"].type == "object"
|
|
assert fields["person.name"].type == "string"
|
|
assert fields["person.tags"].type == "array"
|
|
assert fields["person"].reducer.name == "wf.std.replace"
|
|
assert fields["person.tags"].reducer.name == "wf.std.append"
|
|
|
|
|
|
def test_state_schema_from_preserves_literal_dotted_alias_paths() -> None:
|
|
builder = WorkflowBuilder(
|
|
name="dotted_alias_state_schema_demo",
|
|
input_schema=WorkflowInput,
|
|
state_schema=DotAliasState,
|
|
output_schema=WorkflowOutput,
|
|
start="start",
|
|
)
|
|
|
|
workflow = builder.compile()
|
|
fields = workflow.state_schema.field_index()
|
|
|
|
assert StatePath(("person.name",)) in fields
|
|
assert fields[StatePath(("person.name",))].reducer.name == "wf.std.append"
|