use json schema for state schema

This commit is contained in:
lda
2026-05-20 19:14:00 +07:00 Verified
parent 7322e7ad5f
commit 9f265c3f80
16 changed files with 758 additions and 170 deletions
+230 -31
View File
@@ -1,6 +1,6 @@
from __future__ import annotations
from collections.abc import Mapping
from collections.abc import Iterator, Mapping
from typing import Any
from jsonschema import Draft202012Validator, SchemaError, validators
@@ -121,78 +121,277 @@ class StateFieldDecl(BaseModel):
class StateSchema(BaseModel):
"""Workflow state schema with canonical list fields.
"""Workflow state JSON Schema plus reducer extension keywords.
Deprecated dict-shaped input is still accepted at parse time and normalized
so runtime and serialization only deal with list-of-struct declarations.
Canonical state schemas are ordinary JSON Schema objects. Field-level
workflow metadata such as ``reducer`` and ``trace`` lives beside JSON Schema
keywords inside ``properties`` entries, where JSON Schema validators will
ignore it and wf_core can compile it into runtime behavior.
Deprecated ``fields`` inputs are still accepted at parse time and normalized
into ``properties`` so persisted dumps stay JSON-Schema-shaped.
"""
model_config = ConfigDict(extra="allow")
fields: list[StateFieldDecl] = Field(default_factory=list)
title: str | None = None
type: str | list[str] | None = "object"
properties: dict[str, Any] = Field(default_factory=dict)
required: list[str] = Field(default_factory=list)
@classmethod
def from_field_map(cls, fields: Mapping[str, StateField]) -> StateSchema:
"""Build from the deprecated dict shape at typed Python call sites."""
return cls.model_validate({"fields": fields})
@property
def fields(self) -> list[StateFieldDecl]:
"""Return the compiled field declarations for compatibility callers."""
return list(self.field_map().values())
def field_map(self) -> dict[str, StateFieldDecl]:
"""Return declarations keyed by rootless dotted path."""
return {".".join(field.path.parts): field for field in self.fields}
"""Return reducer-aware declarations keyed by rootless dotted path."""
root_schema = self.model_dump(mode="json", exclude_none=True)
return {
path: field
for path, field in _iter_state_field_declarations(
self.properties,
root_schema,
prefix="",
)
}
def root_fields(self) -> set[str]:
"""Return declared top-level state field names."""
return {field.path.parts[0] for field in self.fields}
return set(self.properties)
@model_serializer(mode="wrap")
def _serialize_without_none_fields(self, handler: Any) -> dict[str, Any]:
"""Persist state schemas as JSON Schema objects without null keywords."""
data = handler(self)
return {key: value for key, value in data.items() if value is not None}
@model_validator(mode="before")
@classmethod
def _coerce_deprecated_field_map(cls, value: object) -> object:
def _coerce_deprecated_fields(cls, value: object) -> object:
if not isinstance(value, Mapping):
return value
data = dict(value)
fields = data.get("fields")
if not isinstance(fields, Mapping):
fields = data.pop("fields", None)
if fields is None:
return data
normalized_fields: list[object] = []
for raw_path, raw_field in fields.items():
path = str(raw_path)
if not path.startswith("state."):
path = f"state.{path}"
if isinstance(fields, list):
for raw_field in fields:
field = StateFieldDecl.model_validate(raw_field)
_set_state_property_schema(
data,
field.path.parts,
_property_schema_from_field(field),
)
return data
if not isinstance(fields, Mapping):
raise ValueError("state_schema.fields must be a mapping or list")
for raw_path, raw_field in fields.items():
if isinstance(raw_field, BaseModel):
field_data = raw_field.model_dump(mode="python")
elif isinstance(raw_field, Mapping):
field_data = dict(raw_field)
if "schema" in field_data or "type" not in field_data:
raise ValueError(
"legacy state field map entries must include 'type'; "
"use canonical list form for entries with 'schema'"
)
else:
raise ValueError(
"legacy state field map entries must include 'type'; "
"use canonical list form for non-legacy declarations"
)
path = str(raw_path)
if not path.startswith("state."):
path = f"state.{path}"
field_data["path"] = path
normalized_fields.append(field_data)
data["fields"] = normalized_fields
field = StateFieldDecl.model_validate(field_data)
_set_state_property_schema(
data,
field.path.parts,
_property_schema_from_field(field),
)
return data
@model_validator(mode="after")
def _reject_duplicate_field_paths(self) -> StateSchema:
seen: set[str] = set()
for field in self.fields:
key = ".".join(field.path.parts)
if key in seen:
raise ValueError(f"duplicate state field path {key!r}")
seen.add(key)
def _validate_state_json_schema_and_extensions(self) -> StateSchema:
schema = self.model_dump(mode="json", exclude_none=True)
validator_cls = (
validators.validator_for(schema)
if "$schema" in schema
else Draft202012Validator
)
try:
validator_cls.check_schema(schema)
except SchemaError as exc:
raise ValueError(f"invalid JSON Schema: {exc.message}") from exc
# JSON Schema permits custom keywords, so wf_core validates reducer
# metadata separately instead of relying on jsonschema to reject it.
for path, property_schema in _iter_property_schemas(
self.properties,
schema,
):
_validate_state_field_extensions(path, property_schema)
return self
def _iter_state_field_declarations(
properties: Mapping[str, Any],
root_schema: Mapping[str, Any],
*,
prefix: str,
) -> Iterator[tuple[str, StateFieldDecl]]:
for name, property_schema in properties.items():
if not isinstance(property_schema, Mapping):
continue
path = f"{prefix}.{name}" if prefix else name
resolved_schema = _resolve_local_ref(property_schema, root_schema)
reducer = _reducer_from_property(path, property_schema)
trace = property_schema.get("trace", True)
default = property_schema.get("default")
if not isinstance(trace, bool):
raise ValueError(f"invalid trace for state field {path!r}: expected bool")
validation_schema = {
key: value
for key, value in resolved_schema.items()
if key not in {"reducer", "trace"}
}
yield (
path,
StateFieldDecl.model_validate(
{
"path": StatePath.of(path),
"schema": SchemaRef.model_validate(validation_schema),
"reducer": reducer,
"trace": trace,
"default": default,
}
),
)
child_properties = resolved_schema.get("properties")
if isinstance(child_properties, Mapping):
yield from _iter_state_field_declarations(
child_properties,
root_schema,
prefix=path,
)
def _iter_property_schemas(
properties: Mapping[str, Any],
root_schema: Mapping[str, Any],
*,
prefix: str = "",
) -> Iterator[tuple[str, Mapping[str, Any]]]:
for name, property_schema in properties.items():
if not isinstance(property_schema, Mapping):
continue
path = f"{prefix}.{name}" if prefix else name
yield path, property_schema
resolved_schema = _resolve_local_ref(property_schema, root_schema)
child_properties = resolved_schema.get("properties")
if isinstance(child_properties, Mapping):
yield from _iter_property_schemas(
child_properties,
root_schema,
prefix=path,
)
def _validate_state_field_extensions(
path: str,
property_schema: Mapping[str, Any],
) -> None:
_reducer_from_property(path, property_schema)
trace = property_schema.get("trace", True)
if not isinstance(trace, bool):
raise ValueError(f"invalid trace for state field {path!r}: expected bool")
def _reducer_from_property(
path: str,
property_schema: Mapping[str, Any],
) -> ReducerRef:
reducer = property_schema.get("reducer", "wf.std.replace")
try:
if isinstance(reducer, str):
return ReducerRef(name=reducer)
if isinstance(reducer, Mapping):
return ReducerRef.model_validate(reducer)
except ValueError as exc:
raise ValueError(f"invalid reducer for state field {path!r}: {exc}") from exc
raise ValueError(
f"invalid reducer for state field {path!r}: expected string or object"
)
def _property_schema_from_field(field: StateFieldDecl) -> dict[str, Any]:
schema = field.validation_schema.model_dump(mode="json", exclude_none=True)
schema["reducer"] = _dump_reducer_keyword(field.reducer)
if not field.trace:
schema["trace"] = False
if field.default is not None:
schema["default"] = field.default
return schema
def _dump_reducer_keyword(reducer: ReducerRef) -> str | dict[str, Any]:
if not reducer.config:
return reducer.name
return reducer.model_dump(mode="json")
def _set_state_property_schema(
data: dict[str, Any],
path_parts: tuple[str, ...],
property_schema: dict[str, Any],
) -> None:
data.setdefault("type", "object")
properties = data.setdefault("properties", {})
if not isinstance(properties, dict):
raise ValueError("state_schema.properties must be an object")
current_properties = properties
for part in path_parts[:-1]:
current = current_properties.setdefault(
part,
{"type": "object", "properties": {}},
)
if not isinstance(current, dict):
raise ValueError(f"state field path {'.'.join(path_parts)!r} overlaps")
current.setdefault("type", "object")
next_properties = current.setdefault("properties", {})
if not isinstance(next_properties, dict):
raise ValueError(f"state field path {'.'.join(path_parts)!r} overlaps")
current_properties = next_properties
leaf = path_parts[-1]
if leaf in current_properties:
raise ValueError(f"duplicate state field path {'.'.join(path_parts)!r}")
current_properties[leaf] = property_schema
def _resolve_local_ref(
property_schema: Mapping[str, Any],
root_schema: Mapping[str, Any],
) -> Mapping[str, Any]:
"""Resolve the common Pydantic ``#/$defs/...`` case for internal indexes."""
ref = property_schema.get("$ref")
if not isinstance(ref, str) or not ref.startswith("#/$defs/"):
return property_schema
definitions = root_schema.get("$defs")
if not isinstance(definitions, Mapping):
return property_schema
resolved = definitions.get(ref.removeprefix("#/$defs/"))
return resolved if isinstance(resolved, Mapping) else property_schema
class NodeDef(BaseModel):
"""Reusable node contract referenced by one or more node uses."""