Files
ww/protocol/check.py

1199 lines
49 KiB
Python

#!/usr/bin/env python3
"""Aggregate WWAR v1 Phase 0 conformance gate."""
from __future__ import annotations
import ast
import copy
import hashlib
import json
import os
import pathlib
import struct
import subprocess
import sys
import tempfile
import unicodedata
from typing import Any
import tool
ROOT = pathlib.Path(__file__).resolve().parent
SCHEMA = ROOT / "schema"
GENERATED = ROOT / "generated" / "wwar_v1_tables.json"
VECTORS = ROOT / "vectors"
EXPECTED_FILES = {
"README.md",
"check.py",
"tool.py",
"schema/digests.json",
"schema/records.json",
"schema/wire.json",
"schema/wwar-v1.sha256",
"generated/wwar_v1_tables.json",
"vectors/digests.json",
"vectors/records.json",
"vectors/wire-invalid.json",
"vectors/wire-valid.json",
}
ARCHIVED_ASSIGNMENT_SHA256 = "7b8660b451cc5260e9c7ad8d45da8e62313ff1f9b54259e211ff04b49c38671f"
def fail(message: str) -> None:
raise AssertionError(message)
def exact_keys(value: Any, expected: set[str], where: str) -> None:
if not isinstance(value, dict) or set(value) != expected:
actual = set(value) if isinstance(value, dict) else set()
fail(
f"{where}: keys differ: missing={sorted(expected-actual)} "
f"unknown={sorted(actual-expected)}"
)
def register_case_id(case: Any, seen: set[str], where: str) -> None:
if not isinstance(case, dict):
fail(f"{where}: expected object")
case_id = case.get("id")
if (
not isinstance(case_id, str)
or not case_id
or len(case_id.encode("utf-8")) > 255
or case_id in seen
):
fail(f"{where}: invalid or duplicate case id")
seen.add(case_id)
def check_files() -> None:
actual = {
path.relative_to(ROOT).as_posix()
for path in ROOT.rglob("*")
if path.is_file()
}
if actual != EXPECTED_FILES:
fail(f"protocol file set differs: missing={sorted(EXPECTED_FILES-actual)} extra={sorted(actual-EXPECTED_FILES)}")
def check_no_bytecode() -> None:
bad = [
path.relative_to(ROOT).as_posix()
for path in ROOT.rglob("*")
if path.is_file() and (path.suffix in {".pyc", ".pyo"} or "__pycache__" in path.parts)
]
if bad:
fail(f"Python bytecode under protocol/: {bad}")
def run_tool(*arguments: str, expect: int = 0) -> subprocess.CompletedProcess[str]:
env = os.environ.copy()
env["PYTHONDONTWRITEBYTECODE"] = "1"
env["LC_ALL"] = "C"
completed = subprocess.run(
[sys.executable, "-B", str(ROOT / "tool.py"), *arguments],
cwd=ROOT,
env=env,
text=True,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
check=False,
)
if completed.returncode != expect:
fail(
f"tool command returned {completed.returncode}, wanted {expect}: {arguments}\n"
f"stdout={completed.stdout}\nstderr={completed.stderr}"
)
return completed
def oracle_encode(ast_value: Any, wire: dict[str, Any], depth: int = 1) -> bytes:
if depth > wire["limits"]["nesting_depth_max"]:
raise ValueError("depth")
if not isinstance(ast_value, dict) or len(ast_value) != 1:
raise ValueError("shape")
name, value = next(iter(ast_value.items()))
codes = {row["name"]: row["code"] for row in wire["wire_types"]}
if name == "bool":
payload = bytes([int(value)])
elif name == "bytes":
payload = bytes.fromhex(value)
elif name == "string":
payload = value.encode("utf-8")
elif name == "uint":
number = int(value)
payload = number.to_bytes(max(1, (number.bit_length() + 7) // 8), "big")
elif name == "list":
children = [oracle_encode(child, wire, depth + 1) for child in value]
payload = struct.pack(">I", len(children)) + b"".join(
struct.pack(">Q", len(child)) + child for child in children
)
elif name == "map":
parts = [struct.pack(">I", len(value))]
for key, child_value in value:
key_bytes = key.encode("utf-8")
child = oracle_encode(child_value, wire, depth + 1)
parts.extend((struct.pack(">Q", len(key_bytes)), key_bytes, struct.pack(">Q", len(child)), child))
payload = b"".join(parts)
elif name == "record":
parts = [struct.pack(">I", len(value))]
for tag, child_value in value:
child = oracle_encode(child_value, wire, depth + 1)
parts.extend((struct.pack(">I", tag), struct.pack(">Q", len(child)), child))
payload = b"".join(parts)
else:
raise ValueError(name)
return bytes([codes[name]]) + struct.pack(">Q", len(payload)) + payload
def oracle_envelope(ast_value: Any, wire: dict[str, Any]) -> bytes:
return bytes.fromhex(wire["magic_hex"]) + struct.pack(">H", wire["schema_version"]) + oracle_encode(ast_value, wire)
def load_vectors(
name: str, expected_format: str, expected_keys: set[str]
) -> dict[str, Any]:
value = tool.load_json(VECTORS / name)
exact_keys(value, expected_keys, name)
if value.get("format") != expected_format or value.get("version") != 1:
fail(f"{name}: bad vector format/version")
return value
def check_wire_vectors(bundle: tool.SchemaBundle) -> None:
valid = load_vectors(
"wire-valid.json", "wwar-v1-wire-valid", {"format", "version", "cases"}
)
ids: set[str] = set()
covered_types: set[str] = set()
for case in valid["cases"]:
if set(case) != {"id", "value", "wwar_hex", "sha256"} or case["id"] in ids:
fail("wire-valid: malformed or duplicate case")
ids.add(case["id"])
encoded = tool.encode_envelope(case["value"], bundle.wire)
oracle = oracle_envelope(case["value"], bundle.wire)
if encoded != oracle or encoded.hex() != case["wwar_hex"]:
fail(f"wire-valid {case['id']}: encoder/oracle/golden mismatch")
if hashlib.sha256(encoded).hexdigest() != case["sha256"]:
fail(f"wire-valid {case['id']}: SHA-256 mismatch")
if tool.decode_envelope(encoded, bundle.wire) != case["value"]:
fail(f"wire-valid {case['id']}: round-trip mismatch")
covered_types.add(next(iter(case["value"])))
if covered_types != set(bundle.type_codes):
fail(f"wire-valid: type coverage differs: {covered_types}")
invalid = load_vectors(
"wire-invalid.json", "wwar-v1-wire-invalid", {"format", "version", "cases"}
)
ids.clear()
required = {
"bytes-limit-before-truncation",
"bytes-truncated",
"bool-truncated-before-length",
"bool-length-mismatch",
"list-child-length-mismatch",
"nesting-limit",
}
for case in invalid["cases"]:
if set(case) != {"id", "input_hex", "error"} or case["id"] in ids:
fail("wire-invalid: malformed or duplicate case")
exact_keys(case["error"], {"code", "path"}, f"wire-invalid {case['id']}.error")
ids.add(case["id"])
try:
tool.decode_envelope(bytes.fromhex(case["input_hex"]), bundle.wire)
except tool.ProtocolError as exc:
if {"code": exc.code, "path": exc.path} != case["error"]:
fail(f"wire-invalid {case['id']}: got {exc.code} {exc.path}, wanted {case['error']}")
else:
fail(f"wire-invalid {case['id']}: accepted")
if not required <= ids:
fail(f"wire-invalid: missing precedence cases {sorted(required-ids)}")
def assignment_catalog(bundle: tool.SchemaBundle) -> dict[str, Any]:
return {
"scalar_types": {
row["name"]: {key: value for key, value in row.items() if key != "name"}
for row in bundle.records_data["scalar_types"]
},
"path_classes": {
row["name"]: {key: value for key, value in row.items() if key != "name"}
for row in bundle.records_data["path_classes"]
},
"records": {
record["name"]: {
"top_level_kind": record["top_level_kind"],
"fields": {
str(field["tag"]): {
key: field[key]
for key in ("name", "type", "cardinality", "encoded_default", "order", "order_by")
}
for field in record["fields"]
},
}
for record in bundle.records_data["records"]
},
"enums": {
enum["name"]: {str(value["value"]): value["name"] for value in enum["values"]}
for enum in bundle.records_data["enums"]
},
"unions": {
union["name"]: {
"record": union["record"],
"discriminator": union["discriminator"],
"cases": {
case["value"]: {
key: sorted(case[key]) for key in ("required", "allowed", "nonempty", "empty")
}
for case in union["cases"]
},
}
for union in bundle.records_data["unions"]
},
"record_kinds": {
str(row["kind"]): {"name": row["name"], "schema": row["schema"]}
for row in bundle.records_data["record_kinds"]
},
"digests": {
row["name"]: {
key: row[key]
for key in ("algorithm", "separator_utf8_hex", "input", "entry_encoding", "formula", "result")
if key in row
}
for row in bundle.digests_data["domains"]
},
"artifact_digest_mapping": {
row["artifact_kind"]: {
key: row[key] for key in ("domain", "record_kind", "record_schema")
}
for row in bundle.records_data["artifact_digest_mapping"]
},
"wrappers": {
row["name"]: {key: row[key] for key in ("body_record", "magic_hex", "identity")}
for row in bundle.records_data["wrappers"]
},
}
def expect_protocol_error(operation: Any, expected: dict[str, str], case_id: str) -> None:
try:
operation()
except tool.ProtocolError as exc:
if {"code": exc.code, "path": exc.path} != expected:
fail(f"{case_id}: got {exc.code} {exc.path}, wanted {expected}")
else:
fail(f"{case_id}: accepted")
def check_record_vectors(bundle: tool.SchemaBundle) -> None:
vectors = load_vectors(
"records.json",
"wwar-v1-record-vectors",
{
"format",
"version",
"archived_assignment_sha256",
"assignment_sha256",
"coverage",
"cases",
"union_cases",
"wrapper_cases",
"invalid_cases",
},
)
if vectors["archived_assignment_sha256"] != ARCHIVED_ASSIGNMENT_SHA256:
fail("records: archived assignment fingerprint changed")
catalog = assignment_catalog(bundle)
digest = hashlib.sha256(
json.dumps(catalog, ensure_ascii=False, sort_keys=True, separators=(",", ":")).encode("utf-8")
).hexdigest()
if digest != vectors["assignment_sha256"]:
fail("records: assignment catalog fingerprint changed")
expected_coverage = {
"scalar_types": len(bundle.scalars),
"path_classes": len(bundle.paths),
"records": len(bundle.records),
"fields": sum(len(row["fields"]) for row in bundle.records.values()),
"encoded_defaults": sum(
field["encoded_default"] is not None
for row in bundle.records.values()
for field in row["fields"]
),
"enums": len(bundle.enums),
"enum_values": sum(len(row["values"]) for row in bundle.enums.values()),
"unions": len(bundle.records_data["unions"]),
"union_arms": sum(len(row["cases"]) for row in bundle.records_data["unions"]),
"record_kinds": len(bundle.kinds),
"artifact_mappings": len(bundle.records_data["artifact_digest_mapping"]),
"wrappers": len(bundle.wrappers),
}
exact_keys(vectors["coverage"], set(expected_coverage), "records.coverage")
if vectors["coverage"] != expected_coverage:
fail(f"records: coverage mismatch {vectors['coverage']} != {expected_coverage}")
for enum in bundle.enums.values():
for member in enum["values"]:
type_spec = {"kind": "enum", "name": enum["name"]}
ast_value = bundle._semantic_ast(type_spec, member["name"], "$.enum")
if ast_value != {"uint": str(member["value"])}:
fail(f"enum {enum['name']}.{member['name']}: typed numeric mismatch")
encoded = tool.encode_envelope(ast_value, bundle.wire)
decoded_ast = tool.decode_envelope(encoded, bundle.wire)
if bundle._decode_semantic_ast(type_spec, decoded_ast, "$.enum") != member["name"]:
fail(f"enum {enum['name']}.{member['name']}: typed round-trip mismatch")
case_ids: set[str] = set()
for case in vectors["cases"]:
allowed = {"id", "record", "value", "wwar_hex", "decoded", "record_id"}
required = {"id", "record", "value", "wwar_hex"}
if not isinstance(case, dict) or not required <= set(case) <= allowed:
fail("records: malformed ordinary case")
register_case_id(case, case_ids, "records ordinary case")
encoded = bundle.encode_record(case["record"], case["value"])
if encoded.hex() != case["wwar_hex"]:
fail(f"record {case['id']}: golden mismatch")
decoded = bundle.decode_record(case["record"], encoded)
if "decoded" in case and decoded != case["decoded"]:
fail(f"record {case['id']}: decoded mismatch")
if "record_id" in case:
kind = bundle.records[case["record"]]["top_level_kind"]
if bundle.record_id(case["record"], kind, 1, case["value"]).hex() != case["record_id"]:
fail(f"record {case['id']}: identity mismatch")
actual_union_cases: set[tuple[str, str]] = set()
for case in vectors["union_cases"]:
exact_keys(case, {"id", "record", "value", "wwar_hex"}, "records.union case")
register_case_id(case, case_ids, "records union case")
encoded = bundle.encode_record(case["record"], case["value"])
if encoded.hex() != case["wwar_hex"] or bundle.decode_record(case["record"], encoded) != case["value"]:
fail(f"union {case['id']}: golden/round-trip mismatch")
actual_union_cases.add((case["record"], case["value"][bundle.unions[case["record"]]["discriminator"]]))
expected_union_cases = {
(union["record"], case["value"])
for union in bundle.records_data["unions"]
for case in union["cases"]
}
if actual_union_cases != expected_union_cases or len(actual_union_cases) != len(vectors["union_cases"]):
fail("union vectors do not cover every arm exactly once")
for case in vectors["wrapper_cases"]:
exact_keys(case, {"id", "wrapper", "value", "bytes_hex"}, "records.wrapper case")
register_case_id(case, case_ids, "records wrapper case")
encoded = bundle.encode_wrapper(case["wrapper"], case["value"])
if encoded.hex() != case["bytes_hex"]:
fail(f"wrapper {case['id']}: golden mismatch")
if bundle.decode_wrapper(case["wrapper"], encoded) != bundle._materialize_record(
bundle.wrappers[case["wrapper"]]["body_record"], case["value"], "$"
):
fail(f"wrapper {case['id']}: round-trip mismatch")
for case in vectors["invalid_cases"]:
register_case_id(case, case_ids, "records invalid case")
if case["operation"] == "record-id":
exact_keys(
case,
{"id", "operation", "record", "kind", "schema", "value", "error"},
"records.invalid record-id",
)
operation = lambda case=case: bundle.record_id(
case["record"], case["kind"], case["schema"], case["value"]
)
elif case["operation"] == "decode-record":
exact_keys(
case,
{"id", "operation", "record", "input_hex", "error"},
"records.invalid decode-record",
)
operation = lambda case=case: bundle.decode_record(
case["record"], bytes.fromhex(case["input_hex"])
)
elif case["operation"] == "decode-wrapper":
exact_keys(
case,
{"id", "operation", "wrapper", "input_hex", "error"},
"records.invalid decode-wrapper",
)
operation = lambda case=case: bundle.decode_wrapper(
case["wrapper"], bytes.fromhex(case["input_hex"])
)
else:
fail(f"record invalid case {case['id']}: unknown operation")
exact_keys(case["error"], {"code", "path"}, f"records.invalid {case['id']}.error")
expect_protocol_error(operation, case["error"], case["id"])
def check_assignment_probes(bundle: tool.SchemaBundle) -> None:
covered_records: set[str] = set()
covered_fields: set[tuple[str, int]] = set()
covered_defaults: set[tuple[str, int]] = set()
covered_kinds: set[int] = set()
for name, record in bundle.records.items():
full = bundle.conformance_sample(name)
encoded = bundle.encode_record(name, full)
raw = tool.decode_envelope(encoded, bundle.wire)
tags = [tag for tag, _ in raw["record"]]
expected_tags = [field["tag"] for field in record["fields"]]
if tags != expected_tags:
fail(f"assignment probe {name}: encoded tags differ")
decoded = bundle.decode_record(name, encoded)
if bundle.encode_record(name, decoded) != encoded:
fail(f"assignment probe {name}: typed codec is not byte-idempotent")
covered_records.add(name)
covered_fields.update((name, tag) for tag in tags)
minimal = {
field["name"]: copy.deepcopy(full[field["name"]])
for field in record["fields"]
if field["encoded_default"] is None
}
union = bundle.unions.get(name)
if union:
selected = full[union["discriminator"]]
case = next(row for row in union["cases"] if row["value"] == selected)
for field_name in case["nonempty"]:
minimal[field_name] = copy.deepcopy(full[field_name])
default_encoded = bundle.encode_record(name, minimal)
default_raw = tool.decode_envelope(default_encoded, bundle.wire)
if [tag for tag, _ in default_raw["record"]] != expected_tags:
fail(f"assignment probe {name}: defaulted tags were omitted")
default_decoded = bundle.decode_record(name, default_encoded)
for field in record["fields"]:
if field["encoded_default"] is None:
continue
if field["name"] not in minimal and default_decoded[field["name"]] != field["encoded_default"]:
fail(f"assignment probe {name}.{field['name']}: encoded default differs")
candidates = union["cases"] if union else [None]
for candidate in candidates:
case_name = candidate["value"] if candidate else None
omitted = bundle.conformance_sample(name, case_name)
omitted.pop(field["name"])
try:
omitted_encoded = bundle.encode_record(name, omitted)
omitted_decoded = bundle.decode_record(name, omitted_encoded)
except (tool.ProtocolError, tool.SchemaError):
continue
omitted_raw = tool.decode_envelope(omitted_encoded, bundle.wire)
if field["tag"] not in {tag for tag, _ in omitted_raw["record"]}:
fail(f"assignment probe {name}.{field['name']}: default tag was omitted")
if omitted_decoded[field["name"]] != field["encoded_default"]:
fail(f"assignment probe {name}.{field['name']}: omitted default differs")
covered_defaults.add((name, field["tag"]))
break
else:
fail(f"assignment probe {name}.{field['name']}: no valid omission probe")
kind = record["top_level_kind"]
if kind is not None:
if len(bundle.record_id(name, kind, 1, full)) != 32:
fail(f"assignment probe {name}: record identity length differs")
covered_kinds.add(kind)
expected_fields = {
(name, field["tag"])
for name, record in bundle.records.items()
for field in record["fields"]
}
expected_defaults = {
(name, field["tag"])
for name, record in bundle.records.items()
for field in record["fields"]
if field["encoded_default"] is not None
}
if covered_records != set(bundle.records):
fail("assignment probes do not cover every record")
if covered_fields != expected_fields:
fail("assignment probes do not cover every field tag")
if covered_defaults != expected_defaults:
fail("assignment probes do not cover every encoded default")
if covered_kinds != set(bundle.kinds):
fail("assignment probes do not cover every record kind")
covered_wrappers = set()
for name, wrapper in bundle.wrappers.items():
value = bundle.conformance_sample(wrapper["body_record"])
encoded = bundle.encode_wrapper(name, value)
decoded = bundle.decode_wrapper(name, encoded)
if bundle.encode_wrapper(name, decoded) != encoded:
fail(f"assignment probe wrapper {name}: codec is not byte-idempotent")
covered_wrappers.add(name)
if covered_wrappers != set(bundle.wrappers):
fail("assignment probes do not cover every wrapper")
order_probe = None
for name, record in bundle.records.items():
for field in record["fields"]:
if field["order"] == "sorted":
order_probe = (name, field)
break
if order_probe:
break
if order_probe is None:
fail("assignment probes found no sorted field")
record_name, field = order_probe
value = bundle.conformance_sample(record_name)
item = bundle._sample_type(field["type"]["item"], (record_name,))
value[field["name"]] = [item, copy.deepcopy(item)]
expect_protocol_error(
lambda: bundle.encode_record(record_name, value),
{"code": "CONSTRAINT_VIOLATION", "path": f"$.{field['name']}[1]"},
"derived-duplicate-sort-key",
)
def oracle_source_tree(entries: list[dict[str, Any]]) -> bytes:
out = bytearray()
for index, entry in enumerate(entries):
exact_keys(
entry,
{"path", "type", "executable", "content"},
f"digests.source-tree entry {index}",
)
path = entry["path"].encode("utf-8")
out += struct.pack(">Q", len(path)) + path
if entry["type"] == "dir":
out += b"\x01\x00" + struct.pack(">Q", 0)
else:
content = bytes.fromhex(entry["content"])
out += b"\x02" + bytes([int(entry["executable"])]) + struct.pack(">Q", len(content)) + content
return bytes(out)
def check_digest_vectors(bundle: tool.SchemaBundle) -> None:
vectors = load_vectors(
"digests.json", "wwar-v1-digest-vectors", {"format", "version", "cases"}
)
if {case["domain"] for case in vectors["cases"]} != set(bundle.domains) or len(vectors["cases"]) != len(bundle.domains):
fail("digests: every domain must have exactly one vector")
action_domains = {
domain["name"]
for domain in bundle.domains.values()
if tuple(domain["formula"]) == tool.FORMULA_WWAR
and set(domain["result"]) == {"kind", "scalar"}
}
if len(action_domains) != 1:
fail("digests: expected one scalar-result WWAR action-key formula")
action_seen: set[str] = set()
case_ids: set[str] = set()
for case in vectors["cases"]:
if not isinstance(case, dict) or "domain" not in case or case["domain"] not in bundle.domains:
fail("digests: malformed or unknown domain case")
register_case_id(case, case_ids, "digests case")
domain = bundle.domains[case["domain"]]
formula = tuple(domain["formula"])
separator = bytes.fromhex(domain["separator_utf8_hex"])
if formula == tool.FORMULA_BYTES:
exact_keys(case, {"id", "domain", "input_hex", "digest"}, "digests.bytes case")
value: Any = bytes.fromhex(case["input_hex"])
expected = hashlib.sha256(separator + struct.pack(">Q", len(value)) + value).digest()
actual = bundle.digest(case["domain"], value)
elif formula == tool.FORMULA_WWAR:
exact_keys(
case,
{"id", "domain", "record", "value", "wwar_hex", "digest"},
"digests.WWAR case",
)
encoded = bundle.encode_record(case["record"], case["value"])
if encoded.hex() != case["wwar_hex"]:
fail(f"digest {case['id']}: WWAR golden mismatch")
expected = hashlib.sha256(separator + struct.pack(">Q", len(encoded)) + encoded).digest()
actual = bundle.digest(case["domain"], case["value"], record_name=case["record"])
if case["domain"] in action_domains:
action_seen.add(case["domain"])
elif formula == tool.FORMULA_RECORD:
exact_keys(
case,
{"id", "domain", "record", "kind", "schema", "value", "wwar_hex", "digest"},
"digests.record case",
)
encoded = bundle.encode_record(case["record"], case["value"])
if encoded.hex() != case["wwar_hex"]:
fail(f"digest {case['id']}: WWAR golden mismatch")
expected = hashlib.sha256(
separator
+ struct.pack(">I", case["kind"])
+ struct.pack(">I", case["schema"])
+ struct.pack(">Q", len(encoded))
+ encoded
).digest()
actual = bundle.digest(
case["domain"],
case["value"],
record_name=case["record"],
kind=case["kind"],
schema=case["schema"],
)
elif formula == tool.FORMULA_SOURCE_TREE:
exact_keys(case, {"id", "domain", "entries", "digest"}, "digests.source-tree case")
value = case["entries"]
expected = hashlib.sha256(separator + oracle_source_tree(value)).digest()
actual = bundle.digest(case["domain"], value)
else:
fail(f"digest {case['id']}: unknown formula")
if actual != expected or actual.hex() != case["digest"]:
fail(f"digest {case['id']}: tool/oracle/golden mismatch")
if action_seen != action_domains:
fail("digests: missing action-key vector")
probe = b"domain separation probe"
separated = []
for domain in bundle.domains.values():
separator = bytes.fromhex(domain["separator_utf8_hex"])
formula = tuple(domain["formula"])
if formula in {tool.FORMULA_BYTES, tool.FORMULA_WWAR}:
preimage = separator + struct.pack(">Q", len(probe)) + probe
elif formula == tool.FORMULA_RECORD:
preimage = separator + struct.pack(">I", 1) + struct.pack(">I", 1) + struct.pack(">Q", len(probe)) + probe
else:
preimage = separator + probe
separated.append(hashlib.sha256(preimage).digest())
if len(separated) != len(set(separated)):
fail("digests: domain separation collision in fixed probe")
def write_schema_set(directory: pathlib.Path, wire: Any, records: Any, digests: Any) -> None:
directory.mkdir(parents=True)
for name, value in (("wire.json", wire), ("records.json", records), ("digests.json", digests)):
(directory / name).write_text(json.dumps(value, ensure_ascii=False, indent=2) + "\n", encoding="utf-8")
tool.write_schema_manifest(directory)
def expect_schema_error(operation: Any, label: str) -> None:
try:
operation()
except tool.SchemaError:
return
fail(f"schema mutation accepted: {label}")
def expect_assertion(operation: Any, label: str) -> None:
try:
operation()
except AssertionError:
return
fail(f"checker mutation accepted: {label}")
def synthetic_records(
record_name: str,
enum_name: str,
field_name: str,
tag: int = 7,
default: str = "red",
) -> dict[str, Any]:
return {
"format": "wwar-v1-record-schema",
"version": 1,
"scalar_types": [],
"path_classes": [],
"enums": [{"name": enum_name, "values": [{"name": "red", "value": 1}, {"name": "blue", "value": 2}]}],
"record_kinds": [{"kind": 1, "name": record_name, "schema": 1}],
"records": [
{
"name": record_name,
"top_level_kind": 1,
"fields": [
{
"tag": tag,
"name": field_name,
"type": {"kind": "enum", "name": enum_name},
"cardinality": "1",
"encoded_default": default,
"order": "none",
"order_by": [],
}
],
}
],
"unions": [],
"artifact_digest_mapping": [],
"wrappers": [],
}
def synthetic_digests(record_name: str) -> dict[str, Any]:
return {
"format": "wwar-v1-digest-schema",
"version": 1,
"domains": [
{
"name": "raw",
"algorithm": "sha256",
"separator_utf8_hex": "53594e2d52415700",
"input": {"kind": "bytes"},
"formula": list(tool.FORMULA_BYTES),
"result": {"kind": "bytes", "length": 32},
},
{
"name": "shape",
"algorithm": "sha256",
"separator_utf8_hex": "53594e2d534841504500",
"input": {"record": record_name, "encoding": "WWAR"},
"formula": list(tool.FORMULA_WWAR),
"result": {"kind": "bytes", "length": 32},
},
{
"name": "identity",
"algorithm": "sha256",
"separator_utf8_hex": "53594e2d4944454e5449545900",
"input": {"record": "top-level record", "encoding": "WWAR"},
"formula": list(tool.FORMULA_RECORD),
"result": {"kind": "bytes", "length": 32},
},
],
}
def replace_exact_strings(value: Any, mapping: dict[str, str]) -> Any:
if isinstance(value, str):
return mapping.get(value, value)
if isinstance(value, list):
return [replace_exact_strings(item, mapping) for item in value]
if isinstance(value, dict):
return {
mapping.get(key, key): replace_exact_strings(child, mapping)
for key, child in value.items()
}
return value
def check_genericity_and_mutations(bundle: tool.SchemaBundle) -> None:
with tempfile.TemporaryDirectory(prefix="wwar-schema-mutations-") as raw_tmp:
base = pathlib.Path(raw_tmp)
wire = copy.deepcopy(bundle.wire)
records = copy.deepcopy(bundle.records_data)
digests = copy.deepcopy(bundle.digests_data)
mutation = copy.deepcopy(records)
mutation["predicate_programs"] = []
directory = base / "unknown-root"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "semantic root")
mutation = copy.deepcopy(wire)
mutation["error_precedence"] = {
"op": "context-lookup",
"relation": "provider-table",
}
directory = base / "wire-policy-object"
write_schema_set(directory, mutation, records, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "wire policy object")
mutation = copy.deepcopy(wire)
mutation["root"] = "X" * 1_000_000
directory = base / "giant-wire-root"
write_schema_set(directory, mutation, records, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "giant wire root")
mutation = copy.deepcopy(records)
mutation["records"][0]["fields"][0]["constraint"] = {"op": "context-lookup"}
directory = base / "field-expression"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "field expression")
mutation = copy.deepcopy(records)
mutation["wrappers"][0]["identity"] = {
"op": "graph-dfs",
"policy": "bootstrap",
}
directory = base / "wrapper-policy-object"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "wrapper policy object")
mutation = copy.deepcopy(records)
mutation["path_classes"][0]["rules"] = [
{"op": "call", "operation": "provider-expand"}
]
directory = base / "path-policy-object"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "path policy object")
mutation = copy.deepcopy(records)
mutation["path_classes"][0]["rules"] = ["no-backslash"]
directory = base / "redundant-path-rules"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "redundant path rules")
mutation = copy.deepcopy(records)
scalar_field = next(
field
for record in mutation["records"]
for field in record["fields"]
if field["type"]["kind"] == "string" and field["cardinality"] == "1"
)
scalar_field["cardinality"] = "*"
directory = base / "scalar-list-cardinality"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "scalar list cardinality")
mutation = copy.deepcopy(records)
scalar_field = next(
field
for record in mutation["records"]
for field in record["fields"]
if field["type"]["kind"] == "string" and field["order"] == "none"
)
scalar_field["order"] = "ordered"
directory = base / "scalar-ordered"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "scalar ordered field")
mutation = copy.deepcopy(records)
scalar_field = next(
field
for record in mutation["records"]
for field in record["fields"]
if field["type"]["kind"] == "string"
and field["cardinality"] == "1"
and field["encoded_default"] is None
and not field["type"].get("scalar")
)
scalar_field["encoded_default"] = "x" * 100_000
directory = base / "giant-string-default"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "giant string default")
mutation = copy.deepcopy(records)
list_field = next(
field
for record in mutation["records"]
for field in record["fields"]
if field["type"]["kind"] == "list"
and field["type"]["item"]["kind"] == "string"
)
list_field["encoded_default"] = ["x"] * 2000
directory = base / "giant-list-default"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "giant list default")
mutation = copy.deepcopy(records)
map_field = next(
field
for record in mutation["records"]
for field in record["fields"]
if field["type"]["kind"] == "map"
)
map_field["encoded_default"] = {f"k{index}": "" for index in range(1100)}
directory = base / "giant-map-default"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "giant map default")
mutation = copy.deepcopy(digests)
mutation["domains"][0]["formula"] = {"op": "call"}
directory = base / "digest-expression"
write_schema_set(directory, wire, records, mutation)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "digest expression")
mutation = copy.deepcopy(digests)
mutation["domains"][0]["result"] = {"scalar": "context-lookup"}
directory = base / "digest-result-policy"
write_schema_set(directory, wire, records, mutation)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "digest result policy")
mutation = copy.deepcopy(records)
incomplete = next(
field
for record in mutation["records"]
for field in record["fields"]
if field["type"]["kind"] == "record"
and isinstance(field["encoded_default"], dict)
and field["encoded_default"]
)
incomplete["encoded_default"].pop(next(iter(incomplete["encoded_default"])))
directory = base / "incomplete-record-default"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "incomplete record default")
mutation = copy.deepcopy(records)
mutation["records"][0]["fields"][1]["tag"] = mutation["records"][0]["fields"][0]["tag"]
directory = base / "duplicate-tag"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "duplicate tag")
mutation = copy.deepcopy(records)
mutation["enums"][0]["values"][1]["value"] = mutation["enums"][0]["values"][0]["value"]
directory = base / "duplicate-enum"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "duplicate enum value")
mutation = copy.deepcopy(digests)
mutation["domains"][1]["separator_utf8_hex"] = mutation["domains"][0]["separator_utf8_hex"]
directory = base / "duplicate-domain"
write_schema_set(directory, wire, records, mutation)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "duplicate digest separator")
mutation = copy.deepcopy(digests)
mutation["domains"][0]["separator_utf8_hex"] = "ff00"
directory = base / "invalid-utf8-separator"
write_schema_set(directory, wire, records, mutation)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "invalid UTF-8 separator")
mutation = copy.deepcopy(digests)
mutation["domains"][0]["separator_utf8_hex"] = "78007900"
directory = base / "interior-nul-separator"
write_schema_set(directory, wire, records, mutation)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "interior NUL separator")
wrong_result_record = next(name for name in bundle.records if name not in bundle.unions)
for formula, label in (
(tool.FORMULA_BYTES, "blob"),
(tool.FORMULA_RECORD, "record"),
):
mutation = copy.deepcopy(digests)
result_domain = next(
domain
for domain in mutation["domains"]
if tuple(domain["formula"]) == formula and "record" in domain["result"]
)
result_domain["result"]["record"] = wrong_result_record
directory = base / f"wrong-{label}-result-record"
write_schema_set(directory, wire, records, mutation)
expect_schema_error(
lambda directory=directory: tool.SchemaBundle.from_dir(directory),
f"wrong {label} result record",
)
mutation = copy.deepcopy(records)
union_case = mutation["unions"][0]["cases"][0]
union_case["allowed"].append(union_case["empty"][0])
directory = base / "union-allowed-empty-overlap"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "union allowed/empty overlap")
mutation = copy.deepcopy(records)
union_case = next(
case
for union in mutation["unions"]
for case in union["cases"]
if case["nonempty"]
)
union_case["required"].remove(union_case["nonempty"][0])
directory = base / "union-nonempty-not-required"
write_schema_set(directory, wire, mutation, digests)
expect_schema_error(lambda: tool.SchemaBundle.from_dir(directory), "union nonempty not required")
duplicate = base / "duplicate.json"
duplicate.write_text('{"format":"x","format":"y"}\n', encoding="utf-8")
expect_schema_error(lambda: tool.load_json(duplicate), "duplicate JSON key")
first_dir = base / "alpha-a"
second_dir = base / "alpha-b"
write_schema_set(first_dir, wire, synthetic_records("Box", "Shade", "hue"), synthetic_digests("Box"))
write_schema_set(second_dir, wire, synthetic_records("Crate", "Tone", "tint"), synthetic_digests("Crate"))
first_table = first_dir / "table.json"
second_table = second_dir / "table.json"
tool.generate_table(first_dir, first_table)
tool.generate_table(second_dir, second_table)
first = tool.load_json(first_table)
second = tool.load_json(second_table)
first.pop("schema_digest")
first.pop("schema_files")
second.pop("schema_digest")
second.pop("schema_files")
second = replace_exact_strings(second, {"Crate": "Box", "Tone": "Shade", "tint": "hue"})
if first != second:
fail("alpha-renamed synthetic schemas changed generator behavior")
third_dir = base / "different-tag-default"
write_schema_set(
third_dir,
wire,
synthetic_records("Vessel", "Hue", "shade", tag=23, default="blue"),
synthetic_digests("Vessel"),
)
third = tool.SchemaBundle.from_dir(third_dir)
encoded = third.encode_record("Vessel", {})
raw = tool.decode_envelope(encoded, third.wire)
if raw["record"][0][0] != 23 or third.decode_record("Vessel", encoded) != {"shade": "blue"}:
fail("synthetic changed tag/default did not drive the generic codec")
expect_assertion(
lambda: exact_keys(
{"format": "x", "version": 1, "cases": [], "provider_policy": {}},
{"format", "version", "cases"},
"synthetic vector",
),
"vector policy root",
)
duplicate_ids: set[str] = set()
register_case_id({"id": "same"}, duplicate_ids, "synthetic vector")
expect_assertion(
lambda: register_case_id({"id": "same"}, duplicate_ids, "synthetic vector"),
"duplicate vector id",
)
stale = tool.load_json(GENERATED)
first_record = next(iter(stale["records"].values()))
if not first_record[1]:
fail("cannot construct stale-output mutation")
first_record[1][0][0] += 100
stale_path = base / "stale.json"
stale_path.write_bytes(tool.canonical_json(stale))
run_tool("check-output", "--schema-dir", str(SCHEMA), "--output", str(stale_path), expect=2)
def check_tool_independence(bundle: tool.SchemaBundle) -> None:
source = (ROOT / "tool.py").read_text(encoding="utf-8")
tree = ast.parse(source, filename="tool.py")
allowed_imports = {
"argparse",
"copy",
"hashlib",
"json",
"pathlib",
"struct",
"sys",
"unicodedata",
"dataclasses",
"typing",
"__future__",
}
forbidden_calls = {"eval", "exec", "compile", "getattr", "setattr", "delattr", "id"}
for node in ast.walk(tree):
if isinstance(node, ast.Import):
for alias in node.names:
if alias.name.split(".")[0] not in allowed_imports:
fail(f"tool.py imports disallowed module {alias.name}")
elif isinstance(node, ast.ImportFrom):
if (node.module or "").split(".")[0] not in allowed_imports:
fail(f"tool.py imports from disallowed module {node.module}")
elif isinstance(node, ast.Call) and isinstance(node.func, ast.Name) and node.func.id in forbidden_calls:
fail(f"tool.py uses forbidden dynamic call {node.func.id}")
action_spellings = {
value["name"]
for enum in bundle.enums.values()
if any("." in value["name"] for value in enum["values"])
for value in enum["values"]
}
semantic_identifiers = (
set(bundle.records)
| set(bundle.enums)
| {row["name"] for row in bundle.records_data["unions"]}
| set(bundle.wrappers)
| (set(bundle.domains) - {"blob", "record"})
| action_spellings
| {
value["name"]
for enum in bundle.enums.values()
for value in enum["values"]
if "." in value["name"]
}
)
for filename in ("tool.py", "check.py"):
parsed = ast.parse((ROOT / filename).read_text(encoding="utf-8"), filename=filename)
for node in ast.walk(parsed):
if isinstance(node, ast.Call) and isinstance(node.func, ast.Name):
if node.func.id in forbidden_calls:
fail(f"{filename} uses forbidden dynamic/object-identity call {node.func.id}")
if (
isinstance(node, ast.Constant)
and isinstance(node.value, str)
and node.value in semantic_identifiers
):
fail(f"{filename} hard-codes schema semantic identifier {node.value}")
def check_deterministic_generation(bundle: tool.SchemaBundle) -> None:
checked = GENERATED.read_bytes()
expected = tool.expected_table(SCHEMA)
if checked != expected:
fail("checked-in generated output is stale")
generated = tool.load_json(GENERATED)
exact_keys(
generated,
{
"format",
"version",
"schema_digest",
"schema_files",
"wire",
"scalars",
"paths",
"enums",
"records",
"unions",
"kinds",
"artifacts",
"wrappers",
"digests",
},
"generated table",
)
if generated != bundle.generated_table():
fail("generated table is not the exact compiled codec projection")
with tempfile.TemporaryDirectory(prefix="wwar-generate-a-") as first_tmp, tempfile.TemporaryDirectory(
prefix="wwar-generate-b-"
) as second_tmp:
first = pathlib.Path(first_tmp) / "generated" / GENERATED.name
second = pathlib.Path(second_tmp) / "generated" / GENERATED.name
run_tool("emit-tables", "--schema-dir", str(SCHEMA), "--output", str(first))
run_tool("emit-tables", "--schema-dir", str(SCHEMA), "--output", str(second))
run_tool("check-output", "--schema-dir", str(SCHEMA), "--output", str(first))
run_tool("check-output", "--schema-dir", str(SCHEMA), "--output", str(second))
if first.read_bytes() != second.read_bytes() or first.read_bytes() != checked:
fail("two clean generations and checked-in output are not byte-identical")
def check_size_budgets() -> None:
byte_limits = {
"README.md": 64 * 1024,
"check.py": 256 * 1024,
"tool.py": 256 * 1024,
"schema/digests.json": 128 * 1024,
"schema/records.json": 1024 * 1024,
"schema/wire.json": 64 * 1024,
"schema/wwar-v1.sha256": 1024,
"generated/wwar_v1_tables.json": 512 * 1024,
"vectors/digests.json": 512 * 1024,
"vectors/records.json": 1024 * 1024,
"vectors/wire-invalid.json": 256 * 1024,
"vectors/wire-valid.json": 256 * 1024,
}
schema_lines = 0
human_lines = 0
human_bytes = 0
total_bytes = 0
for path in sorted(EXPECTED_FILES):
full = ROOT / path
raw = full.read_bytes()
lines = len(raw.splitlines())
total_bytes += len(raw)
if len(raw) > byte_limits[path]:
fail(f"Phase 0 file exceeds byte budget: {path} ({len(raw)} bytes)")
if not path.startswith("generated/"):
human_lines += lines
human_bytes += len(raw)
if path.startswith("schema/") and path.endswith(".json"):
schema_lines += lines
if lines > 5000:
fail(f"human-authored schema file exceeds 5000 lines: {path}")
if schema_lines >= 5000 or human_lines >= 15000 or human_bytes >= 2 * 1024 * 1024:
fail(
"Phase 0 human-authored budget exceeded: "
f"schema_lines={schema_lines}, human_lines={human_lines}, human_bytes={human_bytes}"
)
if total_bytes >= 4 * 1024 * 1024:
fail(f"Phase 0 total byte budget exceeded: {total_bytes}")
def main() -> int:
check_no_bytecode()
check_files()
bundle = tool.SchemaBundle.from_dir(SCHEMA)
check_tool_independence(bundle)
check_deterministic_generation(bundle)
check_wire_vectors(bundle)
check_record_vectors(bundle)
check_assignment_probes(bundle)
check_digest_vectors(bundle)
check_genericity_and_mutations(bundle)
check_size_budgets()
check_no_bytecode()
print(
"WWAR v1 protocol OK: "
f"{len(bundle.records)} records, "
f"{sum(len(row['fields']) for row in bundle.records.values())} fields, "
f"{len(bundle.enums)} enums, {len(bundle.domains)} digest domains"
)
return 0
if __name__ == "__main__":
raise SystemExit(main())