2101 lines
96 KiB
Python
2101 lines
96 KiB
Python
#!/usr/bin/env python3
|
|
"""Strict WWAR v1 schema, codec, table generator, and digest utility."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import argparse
|
|
import copy
|
|
import hashlib
|
|
import json
|
|
import pathlib
|
|
import struct
|
|
import sys
|
|
import unicodedata
|
|
from dataclasses import dataclass
|
|
from typing import Any
|
|
|
|
|
|
SCHEMA_FILES = ("wire.json", "records.json", "digests.json")
|
|
SCHEMA_MANIFEST = "wwar-v1.sha256"
|
|
TABLE_FORMAT = "wwar-v1-tables"
|
|
SCHEMA_DOMAIN = b"WW-SCHEMA-V1\0"
|
|
FORMULA_BYTES = ("separator", "u64be(input-length)", "input-bytes")
|
|
FORMULA_WWAR = ("separator", "u64be(encoded-length)", "encoded-bytes")
|
|
FORMULA_RECORD = (
|
|
"separator",
|
|
"u32be(record-kind)",
|
|
"u32be(record-schema)",
|
|
"u64be(encoded-length)",
|
|
"encoded-bytes",
|
|
)
|
|
FORMULA_SOURCE_TREE = (
|
|
"separator",
|
|
"concatenate entry path,type,executable,content frames",
|
|
)
|
|
WIRE_ROOT = "one complete value; typed identities and wrappers additionally require their declared record"
|
|
WIRE_VALUE_FRAME = {
|
|
"type": "u8",
|
|
"payload_length": "u64be",
|
|
"payload": "exactly payload_length bytes",
|
|
}
|
|
WIRE_CONTAINER_FRAMES = {
|
|
"list": {
|
|
"count": "u32be",
|
|
"members": "count repetitions of u64be(value_length) followed by exactly one complete value frame",
|
|
},
|
|
"map": {
|
|
"count": "u32be",
|
|
"members": "count repetitions of u64be(key_length), raw UTF-8 key, u64be(value_length), complete value frame",
|
|
},
|
|
"record": {
|
|
"count": "u32be",
|
|
"members": "count repetitions of u32be(field_tag), u64be(value_length), complete value frame",
|
|
},
|
|
}
|
|
WIRE_TYPE_ROWS = [
|
|
{"name": "bytes", "code": 1, "payload": {"frame": "raw-bytes", "canonical": ["declared-limit", "frame-available", "exact-length"]}},
|
|
{"name": "string", "code": 2, "payload": {"frame": "utf8", "canonical": ["declared-limit", "frame-available", "valid-utf8", "no-nul-string", "nfc-string", "exact-length"]}},
|
|
{"name": "uint", "code": 3, "payload": {"frame": "minimal-uint", "canonical": ["declared-limit", "frame-available", "uint-width", "uint-minimal", "exact-length"]}},
|
|
{"name": "bool", "code": 4, "payload": {"frame": "byte", "canonical": ["declared-limit", "frame-available", "bool-00-or-01", "exact-length"]}},
|
|
{"name": "list", "code": 5, "payload": {"frame": "counted-values", "canonical": ["declared-limit", "frame-available", "declared-order", "exact-length"]}},
|
|
{"name": "map", "code": 6, "payload": {"frame": "counted-pairs", "canonical": ["declared-limit", "frame-available", "raw-utf8-key-order", "exact-length"]}},
|
|
{"name": "record", "code": 7, "payload": {"frame": "counted-fields", "canonical": ["declared-limit", "frame-available", "increasing-tag-order", "exact-length"]}},
|
|
]
|
|
WIRE_LIMITS = {
|
|
"container_members_max": (1 << 24) - 1,
|
|
"nesting_depth_max": 64,
|
|
"string_or_bytes_length_max": (1 << 31) - 1,
|
|
"uint_max": str((1 << 64) - 1),
|
|
"wwar_payload_length_max": str((1 << 64) - 1),
|
|
}
|
|
WIRE_UNICODE = {
|
|
"version": "16.0.0",
|
|
"form": "NFC",
|
|
"nul_forbidden": True,
|
|
"scalar_values_only": True,
|
|
"utf8": "RFC 3629 shortest form; surrogates and values above U+10FFFF are invalid",
|
|
"validation_order": [
|
|
"decode shortest-form UTF-8",
|
|
"reject NUL and non-scalar values",
|
|
"require exact NFC under Unicode 16.0.0",
|
|
],
|
|
}
|
|
WIRE_ERROR_PRECEDENCE = [
|
|
"declared limit before frame availability",
|
|
"frame truncation before type-specific or exact-length checks",
|
|
"fully framed child length mismatch before child interpretation",
|
|
]
|
|
SCHEMA_LITERAL_BYTES_MAX = 16 * 1024
|
|
SCHEMA_LITERAL_NODES_MAX = 4096
|
|
SCHEMA_LITERAL_MEMBERS_MAX = 1024
|
|
SCHEMA_LITERAL_DEPTH_MAX = 64
|
|
|
|
|
|
class SchemaError(ValueError):
|
|
pass
|
|
|
|
|
|
@dataclass
|
|
class ProtocolError(ValueError):
|
|
code: str
|
|
path: str = "$"
|
|
detail: str = ""
|
|
|
|
def __str__(self) -> str:
|
|
suffix = f": {self.detail}" if self.detail else ""
|
|
return f"{self.code} {self.path}{suffix}"
|
|
|
|
|
|
def _duplicate_object(pairs: list[tuple[str, Any]]) -> dict[str, Any]:
|
|
result: dict[str, Any] = {}
|
|
for key, value in pairs:
|
|
if key in result:
|
|
raise SchemaError(f"duplicate JSON key {key!r}")
|
|
result[key] = value
|
|
return result
|
|
|
|
|
|
def load_json(path: pathlib.Path) -> Any:
|
|
try:
|
|
raw = path.read_bytes()
|
|
except OSError as exc:
|
|
raise SchemaError(f"cannot read {path}: {exc}") from exc
|
|
try:
|
|
text = raw.decode("utf-8")
|
|
except UnicodeDecodeError as exc:
|
|
raise SchemaError(f"{path}: invalid UTF-8") from exc
|
|
if unicodedata.normalize("NFC", text) != text:
|
|
raise SchemaError(f"{path}: JSON text is not NFC")
|
|
try:
|
|
return json.loads(text, object_pairs_hook=_duplicate_object)
|
|
except json.JSONDecodeError as exc:
|
|
raise SchemaError(f"{path}:{exc.lineno}:{exc.colno}: {exc.msg}") from exc
|
|
|
|
|
|
def canonical_json(value: Any) -> bytes:
|
|
return (
|
|
json.dumps(value, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
|
|
+ "\n"
|
|
).encode("utf-8")
|
|
|
|
|
|
def _keys(value: Any, expected: set[str], where: str) -> None:
|
|
if not isinstance(value, dict):
|
|
raise SchemaError(f"{where}: expected object")
|
|
actual = set(value)
|
|
if actual != expected:
|
|
raise SchemaError(
|
|
f"{where}: keys differ: missing={sorted(expected - actual)} "
|
|
f"unknown={sorted(actual - expected)}"
|
|
)
|
|
|
|
|
|
def _allowed_keys(value: Any, allowed: set[str], required: set[str], where: str) -> None:
|
|
if not isinstance(value, dict):
|
|
raise SchemaError(f"{where}: expected object")
|
|
actual = set(value)
|
|
if not required <= actual or not actual <= allowed:
|
|
raise SchemaError(
|
|
f"{where}: keys differ: missing={sorted(required - actual)} "
|
|
f"unknown={sorted(actual - allowed)}"
|
|
)
|
|
|
|
|
|
def _list(value: Any, where: str) -> list[Any]:
|
|
if not isinstance(value, list):
|
|
raise SchemaError(f"{where}: expected array")
|
|
return value
|
|
|
|
|
|
def _name(value: Any, where: str) -> str:
|
|
if (
|
|
not isinstance(value, str)
|
|
or not value
|
|
or len(value.encode("utf-8")) > 255
|
|
or unicodedata.normalize("NFC", value) != value
|
|
):
|
|
raise SchemaError(f"{where}: expected nonempty NFC string")
|
|
return value
|
|
|
|
|
|
def _uint(value: Any, where: str, maximum: int = (1 << 64) - 1) -> int:
|
|
if isinstance(value, bool) or not isinstance(value, int) or value < 0 or value > maximum:
|
|
raise SchemaError(f"{where}: expected unsigned integer <= {maximum}")
|
|
return value
|
|
|
|
|
|
def _hex(value: Any, where: str, exact_bytes: int | None = None) -> bytes:
|
|
if not isinstance(value, str) or len(value) % 2 or value.lower() != value:
|
|
raise SchemaError(f"{where}: expected lowercase even-length hexadecimal")
|
|
try:
|
|
raw = bytes.fromhex(value)
|
|
except ValueError as exc:
|
|
raise SchemaError(f"{where}: invalid hexadecimal") from exc
|
|
if exact_bytes is not None and len(raw) != exact_bytes:
|
|
raise SchemaError(f"{where}: expected {exact_bytes} bytes")
|
|
return raw
|
|
|
|
|
|
def _validate_schema_literal_budget(value: Any, where: str) -> None:
|
|
nodes = 0
|
|
|
|
def walk(child: Any, depth: int) -> None:
|
|
nonlocal nodes
|
|
nodes += 1
|
|
if nodes > SCHEMA_LITERAL_NODES_MAX:
|
|
raise SchemaError(f"{where}: schema literal has too many nodes")
|
|
if depth > SCHEMA_LITERAL_DEPTH_MAX:
|
|
raise SchemaError(f"{where}: schema literal is too deeply nested")
|
|
if isinstance(child, str):
|
|
if len(child.encode("utf-8")) > SCHEMA_LITERAL_BYTES_MAX:
|
|
raise SchemaError(f"{where}: schema literal string is too large")
|
|
return
|
|
if isinstance(child, list):
|
|
if len(child) > SCHEMA_LITERAL_MEMBERS_MAX:
|
|
raise SchemaError(f"{where}: schema literal list has too many members")
|
|
for item in child:
|
|
walk(item, depth + 1)
|
|
return
|
|
if isinstance(child, dict):
|
|
if len(child) > SCHEMA_LITERAL_MEMBERS_MAX:
|
|
raise SchemaError(f"{where}: schema literal map has too many members")
|
|
for key, item in child.items():
|
|
walk(key, depth + 1)
|
|
walk(item, depth + 1)
|
|
|
|
walk(value, 0)
|
|
try:
|
|
encoded = canonical_json(value)
|
|
except (TypeError, ValueError) as exc:
|
|
raise SchemaError(f"{where}: schema literal is not JSON data") from exc
|
|
if len(encoded) > SCHEMA_LITERAL_BYTES_MAX:
|
|
raise SchemaError(f"{where}: canonical schema literal is too large")
|
|
|
|
|
|
def _compiled_type(value: dict[str, Any]) -> list[Any]:
|
|
kind = value["kind"]
|
|
if kind in {"bool", "bytes", "string", "uint"}:
|
|
return [kind, value.get("scalar")]
|
|
if kind in {"enum", "record"}:
|
|
return [kind, value["name"]]
|
|
if kind == "list":
|
|
return [kind, _compiled_type(value["item"])]
|
|
if kind == "map":
|
|
return [kind, _compiled_type(value["value"])]
|
|
raise SchemaError(f"cannot compile unknown type {kind!r}")
|
|
|
|
|
|
class SchemaBundle:
|
|
def __init__(
|
|
self,
|
|
wire: dict[str, Any],
|
|
records: dict[str, Any],
|
|
digests: dict[str, Any],
|
|
file_hashes: dict[str, str] | None = None,
|
|
schema_digest: str | None = None,
|
|
) -> None:
|
|
self.wire = wire
|
|
self.records_data = records
|
|
self.digests_data = digests
|
|
self.file_hashes = file_hashes or {}
|
|
self.schema_digest = schema_digest or ""
|
|
self.records: dict[str, dict[str, Any]] = {}
|
|
self.enums: dict[str, dict[str, Any]] = {}
|
|
self.scalars: dict[str, dict[str, Any]] = {}
|
|
self.paths: dict[str, dict[str, Any]] = {}
|
|
self.unions: dict[str, dict[str, Any]] = {}
|
|
self.kinds: dict[int, dict[str, Any]] = {}
|
|
self.domains: dict[str, dict[str, Any]] = {}
|
|
self.wrappers: dict[str, dict[str, Any]] = {}
|
|
self.type_codes: dict[str, int] = {}
|
|
self.type_names: dict[int, str] = {}
|
|
|
|
@classmethod
|
|
def from_dir(cls, schema_dir: pathlib.Path, check_manifest: bool = True) -> "SchemaBundle":
|
|
paths = {name: schema_dir / name for name in SCHEMA_FILES}
|
|
values = {name: load_json(path) for name, path in paths.items()}
|
|
file_hashes = {
|
|
name: hashlib.sha256(paths[name].read_bytes()).hexdigest() for name in SCHEMA_FILES
|
|
}
|
|
digest = schema_bundle_digest(paths)
|
|
bundle = cls(
|
|
values["wire.json"],
|
|
values["records.json"],
|
|
values["digests.json"],
|
|
file_hashes,
|
|
digest,
|
|
)
|
|
bundle.validate()
|
|
if check_manifest:
|
|
check_schema_manifest(schema_dir, file_hashes)
|
|
return bundle
|
|
|
|
def validate(self) -> None:
|
|
self._validate_wire()
|
|
self._validate_records()
|
|
self._validate_digests()
|
|
|
|
def _validate_wire(self) -> None:
|
|
_keys(
|
|
self.wire,
|
|
{
|
|
"format",
|
|
"version",
|
|
"magic_hex",
|
|
"schema_version",
|
|
"root",
|
|
"value_frame",
|
|
"container_frames",
|
|
"wire_types",
|
|
"limits",
|
|
"unicode",
|
|
"error_precedence",
|
|
},
|
|
"$wire",
|
|
)
|
|
if self.wire["format"] != "wwar-v1-wire-schema" or self.wire["version"] != 1:
|
|
raise SchemaError("$wire: unsupported format/version")
|
|
_hex(self.wire["magic_hex"], "$wire.magic_hex", 4)
|
|
if self.wire["magic_hex"] != "57574152":
|
|
raise SchemaError("$wire.magic_hex: WWAR v1 magic differs")
|
|
_uint(self.wire["schema_version"], "$wire.schema_version", 0xFFFF)
|
|
if self.wire["schema_version"] != 1:
|
|
raise SchemaError("$wire.schema_version: WWAR v1 schema differs")
|
|
if self.wire["root"] != WIRE_ROOT:
|
|
raise SchemaError("$wire.root: unsupported root framing")
|
|
_keys(self.wire["value_frame"], {"type", "payload_length", "payload"}, "$wire.value_frame")
|
|
if self.wire["value_frame"] != WIRE_VALUE_FRAME:
|
|
raise SchemaError("$wire.value_frame: unsupported framing")
|
|
containers = self.wire["container_frames"]
|
|
_keys(containers, {"list", "map", "record"}, "$wire.container_frames")
|
|
_keys(containers["list"], {"count", "members"}, "$wire.container_frames.list")
|
|
_keys(containers["map"], {"count", "members"}, "$wire.container_frames.map")
|
|
_keys(containers["record"], {"count", "members"}, "$wire.container_frames.record")
|
|
if containers != WIRE_CONTAINER_FRAMES:
|
|
raise SchemaError("$wire.container_frames: unsupported framing")
|
|
seen_names: set[str] = set()
|
|
seen_codes: set[int] = set()
|
|
for index, row in enumerate(_list(self.wire["wire_types"], "$wire.wire_types")):
|
|
_keys(row, {"name", "code", "payload"}, f"$wire.wire_types[{index}]")
|
|
name = _name(row["name"], f"$wire.wire_types[{index}].name")
|
|
code = _uint(row["code"], f"$wire.wire_types[{index}].code", 255)
|
|
_keys(row["payload"], {"frame", "canonical"}, f"$wire.wire_types[{index}].payload")
|
|
if name in seen_names or code in seen_codes:
|
|
raise SchemaError("$wire.wire_types: duplicate name or code")
|
|
seen_names.add(name)
|
|
seen_codes.add(code)
|
|
if self.wire["wire_types"] != WIRE_TYPE_ROWS:
|
|
raise SchemaError("$wire.wire_types: WWAR v1 type table differs")
|
|
self.type_codes = {row["name"]: row["code"] for row in self.wire["wire_types"]}
|
|
self.type_names = {code: name for name, code in self.type_codes.items()}
|
|
_keys(
|
|
self.wire["limits"],
|
|
{
|
|
"container_members_max",
|
|
"nesting_depth_max",
|
|
"string_or_bytes_length_max",
|
|
"uint_max",
|
|
"wwar_payload_length_max",
|
|
},
|
|
"$wire.limits",
|
|
)
|
|
if self.wire["limits"] != WIRE_LIMITS:
|
|
raise SchemaError("$wire.limits: WWAR v1 limits differ")
|
|
_keys(
|
|
self.wire["unicode"],
|
|
{"version", "form", "nul_forbidden", "scalar_values_only", "utf8", "validation_order"},
|
|
"$wire.unicode",
|
|
)
|
|
if self.wire["unicode"] != WIRE_UNICODE:
|
|
raise SchemaError("$wire.unicode: WWAR v1 Unicode rules differ")
|
|
if (
|
|
self.wire["unicode"]["nul_forbidden"] is not True
|
|
or self.wire["unicode"]["scalar_values_only"] is not True
|
|
):
|
|
raise SchemaError("$wire.unicode: canonical flags must be booleans")
|
|
if WIRE_UNICODE["version"] != unicodedata.unidata_version:
|
|
raise SchemaError(
|
|
f"$wire.unicode.version: need {WIRE_UNICODE['version']}, "
|
|
f"runtime has {unicodedata.unidata_version}"
|
|
)
|
|
if self.wire["error_precedence"] != WIRE_ERROR_PRECEDENCE:
|
|
raise SchemaError("$wire.error_precedence: WWAR v1 precedence differs")
|
|
|
|
def _validate_records(self) -> None:
|
|
data = self.records_data
|
|
_keys(
|
|
data,
|
|
{
|
|
"format",
|
|
"version",
|
|
"scalar_types",
|
|
"path_classes",
|
|
"enums",
|
|
"record_kinds",
|
|
"records",
|
|
"unions",
|
|
"artifact_digest_mapping",
|
|
"wrappers",
|
|
},
|
|
"$records",
|
|
)
|
|
if data["format"] != "wwar-v1-record-schema" or data["version"] != 1:
|
|
raise SchemaError("$records: unsupported format/version")
|
|
scalar_names: set[str] = set()
|
|
for index, scalar in enumerate(_list(data["scalar_types"], "$records.scalar_types")):
|
|
where = f"$records.scalar_types[{index}]"
|
|
_allowed_keys(
|
|
scalar,
|
|
{"name", "base", "length", "maximum", "values", "path_class", "nonempty"},
|
|
{"name", "base"},
|
|
where,
|
|
)
|
|
name = _name(scalar["name"], where + ".name")
|
|
if name in scalar_names:
|
|
raise SchemaError(f"{where}: duplicate scalar {name}")
|
|
scalar_names.add(name)
|
|
if scalar["base"] not in {"bool", "bytes", "string", "uint"}:
|
|
raise SchemaError(f"{where}.base: invalid primitive")
|
|
if "length" in scalar:
|
|
length = _uint(scalar["length"], where + ".length")
|
|
if scalar["base"] != "bytes" or length == 0:
|
|
raise SchemaError(f"{where}.length: only nonempty byte lengths are supported")
|
|
if "maximum" in scalar:
|
|
maximum = _uint(scalar["maximum"], where + ".maximum")
|
|
if scalar["base"] != "uint" or maximum == 0:
|
|
raise SchemaError(f"{where}.maximum: only positive uint maxima are supported")
|
|
if "values" in scalar:
|
|
values = _list(scalar["values"], where + ".values")
|
|
_validate_schema_literal_budget(values, where + ".values")
|
|
if not values or len(values) != len({json.dumps(item, sort_keys=True) for item in values}):
|
|
raise SchemaError(f"{where}.values: expected nonempty unique values")
|
|
for value_index, value in enumerate(values):
|
|
value_where = f"{where}.values[{value_index}]"
|
|
if scalar["base"] == "string":
|
|
_name(value, value_where) if value else _validate_string(value, value_where)
|
|
elif scalar["base"] == "uint":
|
|
_uint(value, value_where)
|
|
elif scalar["base"] == "bool" and not isinstance(value, bool):
|
|
raise SchemaError(f"{value_where}: expected bool")
|
|
elif scalar["base"] == "bytes":
|
|
_hex(value, value_where)
|
|
if "nonempty" in scalar and scalar["nonempty"] is not True:
|
|
raise SchemaError(f"{where}.nonempty: expected true")
|
|
path_names: set[str] = set()
|
|
path_allowed = {
|
|
"name",
|
|
"shape",
|
|
"absolute",
|
|
"allow_complete_dot",
|
|
"allow_dotdot_segments",
|
|
"ascii",
|
|
"roots",
|
|
}
|
|
for index, path in enumerate(_list(data["path_classes"], "$records.path_classes")):
|
|
where = f"$records.path_classes[{index}]"
|
|
_allowed_keys(path, path_allowed, path_allowed - {"roots"}, where)
|
|
name = _name(path["name"], where + ".name")
|
|
if name in path_names:
|
|
raise SchemaError(f"{where}: duplicate path class {name}")
|
|
path_names.add(name)
|
|
if path["shape"] not in {"path", "single-segment"}:
|
|
raise SchemaError(f"{where}.shape: invalid")
|
|
if not (
|
|
isinstance(path["absolute"], bool) or path["absolute"] == "either"
|
|
):
|
|
raise SchemaError(f"{where}.absolute: invalid")
|
|
for flag in ("allow_complete_dot", "allow_dotdot_segments", "ascii"):
|
|
if not isinstance(path[flag], bool):
|
|
raise SchemaError(f"{where}.{flag}: expected bool")
|
|
if "roots" in path:
|
|
roots = _list(path["roots"], where + ".roots")
|
|
if (
|
|
path["absolute"] is not True
|
|
or not roots
|
|
or len(roots) != len(set(roots))
|
|
or not all(
|
|
isinstance(root, str)
|
|
and root.startswith("/")
|
|
and len(root.encode("utf-8")) <= 255
|
|
and unicodedata.normalize("NFC", root) == root
|
|
and "//" not in root
|
|
and "\x00" not in root
|
|
for root in roots
|
|
)
|
|
):
|
|
raise SchemaError(f"{where}.roots: invalid virtual roots")
|
|
for scalar in data["scalar_types"]:
|
|
if "path_class" in scalar and scalar["path_class"] not in path_names:
|
|
raise SchemaError(f"$records.scalar_types[{scalar['name']}]: unknown path class")
|
|
|
|
enum_names: set[str] = set()
|
|
for index, enum in enumerate(_list(data["enums"], "$records.enums")):
|
|
where = f"$records.enums[{index}]"
|
|
_keys(enum, {"name", "values"}, where)
|
|
name = _name(enum["name"], where + ".name")
|
|
if name in enum_names:
|
|
raise SchemaError(f"{where}: duplicate enum {name}")
|
|
enum_names.add(name)
|
|
names: set[str] = set()
|
|
values: set[int] = set()
|
|
previous = -1
|
|
for value_index, value in enumerate(_list(enum["values"], where + ".values")):
|
|
value_where = f"{where}.values[{value_index}]"
|
|
_keys(value, {"name", "value"}, value_where)
|
|
value_name = _name(value["name"], value_where + ".name")
|
|
number = _uint(value["value"], value_where + ".value")
|
|
if value_name in names or number in values:
|
|
raise SchemaError(f"{value_where}: duplicate enum name/value")
|
|
if number <= previous:
|
|
raise SchemaError(f"{value_where}: enum values must increase")
|
|
previous = number
|
|
names.add(value_name)
|
|
values.add(number)
|
|
|
|
record_names: set[str] = set()
|
|
for index, record in enumerate(_list(data["records"], "$records.records")):
|
|
where = f"$records.records[{index}]"
|
|
_keys(record, {"name", "top_level_kind", "fields"}, where)
|
|
name = _name(record["name"], where + ".name")
|
|
if name in record_names:
|
|
raise SchemaError(f"{where}: duplicate record {name}")
|
|
record_names.add(name)
|
|
if record["top_level_kind"] is not None:
|
|
_uint(record["top_level_kind"], where + ".top_level_kind", 0xFFFFFFFF)
|
|
previous = 0
|
|
field_names: set[str] = set()
|
|
for field_index, field in enumerate(_list(record["fields"], where + ".fields")):
|
|
field_where = f"{where}.fields[{field_index}]"
|
|
_keys(
|
|
field,
|
|
{"tag", "name", "type", "cardinality", "encoded_default", "order", "order_by"},
|
|
field_where,
|
|
)
|
|
tag = _uint(field["tag"], field_where + ".tag", 0xFFFFFFFF)
|
|
field_name = _name(field["name"], field_where + ".name")
|
|
if tag <= previous or field_name in field_names:
|
|
raise SchemaError(f"{field_where}: tags must increase and names be unique")
|
|
previous = tag
|
|
field_names.add(field_name)
|
|
if field["cardinality"] not in {"0/1", "1", "*", "map"}:
|
|
raise SchemaError(f"{field_where}.cardinality: invalid")
|
|
if field["order"] not in {"none", "ordered", "sorted", "raw-utf8"}:
|
|
raise SchemaError(f"{field_where}.order: invalid")
|
|
if not isinstance(field["order_by"], list) or not all(
|
|
isinstance(item, str) and item for item in field["order_by"]
|
|
):
|
|
raise SchemaError(f"{field_where}.order_by: invalid")
|
|
self.records = {row["name"]: row for row in data["records"]}
|
|
self.enums = {row["name"]: row for row in data["enums"]}
|
|
self.scalars = {row["name"]: row for row in data["scalar_types"]}
|
|
self.paths = {row["name"]: row for row in data["path_classes"]}
|
|
for record_index, record in enumerate(data["records"]):
|
|
for field_index, field in enumerate(record["fields"]):
|
|
where = f"$records.records[{record_index}].fields[{field_index}]"
|
|
self._validate_type(field["type"], where + ".type")
|
|
self._validate_order_spec(field, where)
|
|
if field["encoded_default"] is not None:
|
|
_validate_schema_literal_budget(
|
|
field["encoded_default"], where + ".encoded_default"
|
|
)
|
|
self._validate_semantic_value(
|
|
field["type"],
|
|
field["encoded_default"],
|
|
where + ".encoded_default",
|
|
allow_empty_record=field["cardinality"] == "0/1",
|
|
allow_empty_scalar=field["cardinality"] == "0/1",
|
|
)
|
|
|
|
kind_numbers: set[int] = set()
|
|
kind_records: set[str] = set()
|
|
previous_kind = 0
|
|
for index, row in enumerate(_list(data["record_kinds"], "$records.record_kinds")):
|
|
where = f"$records.record_kinds[{index}]"
|
|
_keys(row, {"kind", "name", "schema"}, where)
|
|
kind = _uint(row["kind"], where + ".kind", 0xFFFFFFFF)
|
|
name = _name(row["name"], where + ".name")
|
|
schema = _uint(row["schema"], where + ".schema", 0xFFFFFFFF)
|
|
if kind <= previous_kind or kind in kind_numbers or name in kind_records:
|
|
raise SchemaError(f"{where}: kinds must increase and be unique")
|
|
previous_kind = kind
|
|
kind_numbers.add(kind)
|
|
kind_records.add(name)
|
|
if name not in self.records or self.records[name]["top_level_kind"] != kind or schema != 1:
|
|
raise SchemaError(f"{where}: kind/record/schema mismatch")
|
|
for record in data["records"]:
|
|
if record["top_level_kind"] is not None and record["name"] not in kind_records:
|
|
raise SchemaError(f"$records.records[{record['name']}]: missing record-kind row")
|
|
self.kinds = {row["kind"]: row for row in data["record_kinds"]}
|
|
|
|
union_names: set[str] = set()
|
|
union_records: set[str] = set()
|
|
for index, union in enumerate(_list(data["unions"], "$records.unions")):
|
|
where = f"$records.unions[{index}]"
|
|
_keys(union, {"name", "record", "discriminator", "cases"}, where)
|
|
union_name = _name(union["name"], where + ".name")
|
|
record_name = _name(union["record"], where + ".record")
|
|
discriminator = _name(union["discriminator"], where + ".discriminator")
|
|
if union_name in union_names or record_name in union_records or record_name not in self.records:
|
|
raise SchemaError(f"{where}: duplicate or unknown union record")
|
|
union_names.add(union_name)
|
|
union_records.add(record_name)
|
|
fields = {field["name"]: field for field in self.records[record_name]["fields"]}
|
|
if discriminator not in fields or fields[discriminator]["type"].get("kind") != "enum":
|
|
raise SchemaError(f"{where}: discriminator is not an enum field")
|
|
enum = self.enums[fields[discriminator]["type"]["name"]]
|
|
enum_members = {row["name"] for row in enum["values"]}
|
|
cases: set[str] = set()
|
|
for case_index, case in enumerate(_list(union["cases"], where + ".cases")):
|
|
case_where = f"{where}.cases[{case_index}]"
|
|
_keys(case, {"value", "required", "allowed", "nonempty", "empty"}, case_where)
|
|
value = _name(case["value"], case_where + ".value")
|
|
if value in cases or value not in enum_members:
|
|
raise SchemaError(f"{case_where}: duplicate or unknown discriminator value")
|
|
cases.add(value)
|
|
groups: dict[str, set[str]] = {}
|
|
for group in ("required", "allowed", "nonempty", "empty"):
|
|
values = _list(case[group], case_where + "." + group)
|
|
if not all(isinstance(item, str) and item in fields for item in values):
|
|
raise SchemaError(f"{case_where}.{group}: unknown field")
|
|
if len(values) != len(set(values)):
|
|
raise SchemaError(f"{case_where}.{group}: duplicate field")
|
|
groups[group] = set(values)
|
|
if groups["allowed"] & groups["empty"]:
|
|
raise SchemaError(f"{case_where}: allowed/empty fields must be disjoint")
|
|
if groups["allowed"] | groups["empty"] != set(fields):
|
|
raise SchemaError(f"{case_where}: allowed/empty fields do not partition record")
|
|
if not groups["required"] <= groups["allowed"]:
|
|
raise SchemaError(f"{case_where}: required field not allowed")
|
|
if not groups["nonempty"] <= groups["required"]:
|
|
raise SchemaError(f"{case_where}: nonempty field must be required")
|
|
if discriminator not in groups["required"]:
|
|
raise SchemaError(f"{case_where}: discriminator must be required")
|
|
for field_name in set(fields) - groups["required"]:
|
|
if fields[field_name]["encoded_default"] is None:
|
|
raise SchemaError(
|
|
f"{case_where}: non-required field {field_name!r} needs a default"
|
|
)
|
|
for field_name in groups["empty"]:
|
|
default = fields[field_name]["encoded_default"]
|
|
if default is None or not _is_empty(default):
|
|
raise SchemaError(
|
|
f"{case_where}: inactive field {field_name!r} needs an empty default"
|
|
)
|
|
if cases != enum_members:
|
|
raise SchemaError(f"{where}: union cases do not cover discriminator enum")
|
|
self.unions = {row["record"]: row for row in data["unions"]}
|
|
|
|
mapped_artifacts: set[str] = set()
|
|
for index, row in enumerate(
|
|
_list(data["artifact_digest_mapping"], "$records.artifact_digest_mapping")
|
|
):
|
|
where = f"$records.artifact_digest_mapping[{index}]"
|
|
_keys(row, {"artifact_kind", "domain", "record_kind", "record_schema"}, where)
|
|
artifact = _name(row["artifact_kind"], where + ".artifact_kind")
|
|
if artifact in mapped_artifacts:
|
|
raise SchemaError(f"{where}: duplicate artifact kind")
|
|
mapped_artifacts.add(artifact)
|
|
if row["domain"] not in {"blob", "record"}:
|
|
raise SchemaError(f"{where}.domain: invalid")
|
|
kind = _uint(row["record_kind"], where + ".record_kind", 0xFFFFFFFF)
|
|
schema = _uint(row["record_schema"], where + ".record_schema", 0xFFFFFFFF)
|
|
if row["domain"] == "blob" and (kind, schema) != (0, 0):
|
|
raise SchemaError(f"{where}: blob mapping must use kind/schema zero")
|
|
if row["domain"] == "record" and (kind not in kind_numbers or schema != 1):
|
|
raise SchemaError(f"{where}: unknown record identity")
|
|
if mapped_artifacts and not any(
|
|
{value["name"] for value in enum["values"]} == mapped_artifacts
|
|
for enum in data["enums"]
|
|
):
|
|
raise SchemaError("$records.artifact_digest_mapping: no matching closed enum")
|
|
|
|
wrapper_names: set[str] = set()
|
|
for index, row in enumerate(_list(data["wrappers"], "$records.wrappers")):
|
|
where = f"$records.wrappers[{index}]"
|
|
_allowed_keys(
|
|
row,
|
|
{"name", "magic_hex", "framing", "trailing_bytes", "body_record", "identity"},
|
|
{"name", "magic_hex", "body_record", "identity"},
|
|
where,
|
|
)
|
|
name = _name(row["name"], where + ".name")
|
|
if name in wrapper_names or row["body_record"] not in self.records:
|
|
raise SchemaError(f"{where}: duplicate wrapper or unknown body record")
|
|
wrapper_names.add(name)
|
|
magic = _hex(row["magic_hex"], where + ".magic_hex")
|
|
identity = row["identity"]
|
|
kind = self.records[row["body_record"]]["top_level_kind"]
|
|
if kind is None:
|
|
raise SchemaError(f"{where}.body_record: wrapper body needs a record kind")
|
|
expected_identity = f"record_id({kind},1,body)"
|
|
if "framing" in row:
|
|
expected_identity += (
|
|
"; the exact eight-byte magic is verified and reconstructed framing "
|
|
"and is not a second identity"
|
|
)
|
|
if identity != expected_identity:
|
|
raise SchemaError(f"{where}.identity: wrapper identity differs")
|
|
if "trailing_bytes" in row and row["trailing_bytes"] != "reject":
|
|
raise SchemaError(f"{where}.trailing_bytes: only reject is supported")
|
|
if "framing" in row:
|
|
frames = _list(row["framing"], where + ".framing")
|
|
if not frames:
|
|
raise SchemaError(f"{where}.framing: expected at least one frame")
|
|
for frame_index, frame in enumerate(frames):
|
|
_allowed_keys(
|
|
frame,
|
|
{"offset", "length", "value_hex", "encoding"},
|
|
{"offset", "length"},
|
|
f"{where}.framing[{frame_index}]",
|
|
)
|
|
frame_where = f"{where}.framing[{frame_index}]"
|
|
_uint(frame["offset"], frame_where + ".offset")
|
|
length = frame["length"]
|
|
if length != "to EOF":
|
|
_uint(length, frame_where + ".length")
|
|
has_value = "value_hex" in frame
|
|
has_encoding = "encoding" in frame
|
|
if has_value == has_encoding:
|
|
raise SchemaError(f"{frame_where}: need exactly one value or encoding")
|
|
if has_value:
|
|
raw = _hex(frame["value_hex"], frame_where + ".value_hex")
|
|
if not isinstance(length, int) or len(raw) != length:
|
|
raise SchemaError(f"{frame_where}: literal length mismatch")
|
|
else:
|
|
encoding = frame["encoding"]
|
|
if (
|
|
not isinstance(encoding, str)
|
|
or not encoding
|
|
or len(encoding.encode("utf-8")) > 512
|
|
or unicodedata.normalize("NFC", encoding) != encoding
|
|
):
|
|
raise SchemaError(f"{frame_where}.encoding: invalid")
|
|
expected_frames = [
|
|
{"offset": 0, "length": len(magic), "value_hex": magic.hex()},
|
|
{
|
|
"offset": len(magic),
|
|
"length": "to EOF",
|
|
"encoding": f"one complete WWAR version-1 {row['body_record']} envelope",
|
|
},
|
|
]
|
|
if frames != expected_frames or row.get("trailing_bytes") != "reject":
|
|
raise SchemaError(f"{where}.framing: wrapper framing differs")
|
|
elif magic or "trailing_bytes" in row:
|
|
raise SchemaError(f"{where}: magic wrappers require explicit framing")
|
|
self.wrappers = {row["name"]: row for row in data["wrappers"]}
|
|
|
|
def _validate_type(self, value: Any, where: str) -> None:
|
|
if not isinstance(value, dict) or "kind" not in value:
|
|
raise SchemaError(f"{where}: invalid type")
|
|
kind = value["kind"]
|
|
if kind in {"bool", "bytes", "string", "uint"}:
|
|
_allowed_keys(value, {"kind", "scalar"}, {"kind"}, where)
|
|
if "scalar" in value:
|
|
scalar = value["scalar"]
|
|
if scalar not in self.scalars or self.scalars[scalar]["base"] != kind:
|
|
raise SchemaError(f"{where}: scalar base mismatch")
|
|
elif kind in {"enum", "record"}:
|
|
_keys(value, {"kind", "name"}, where)
|
|
catalog = self.enums if kind == "enum" else self.records
|
|
if value["name"] not in catalog:
|
|
raise SchemaError(f"{where}: unknown {kind} {value['name']!r}")
|
|
elif kind == "list":
|
|
_keys(value, {"kind", "item"}, where)
|
|
self._validate_type(value["item"], where + ".item")
|
|
elif kind == "map":
|
|
_keys(value, {"kind", "value"}, where)
|
|
self._validate_type(value["value"], where + ".value")
|
|
else:
|
|
raise SchemaError(f"{where}: unknown type kind {kind!r}")
|
|
|
|
def _validate_order_spec(self, field: dict[str, Any], where: str) -> None:
|
|
order = field["order"]
|
|
order_by = field["order_by"]
|
|
type_spec = field["type"]
|
|
kind = type_spec["kind"]
|
|
cardinality = field["cardinality"]
|
|
if kind == "list":
|
|
if cardinality != "*" or order not in {"ordered", "sorted"}:
|
|
raise SchemaError(f"{where}: lists require * and ordered/sorted")
|
|
elif kind == "map":
|
|
if cardinality != "map" or order != "raw-utf8":
|
|
raise SchemaError(f"{where}: maps require map and raw-utf8")
|
|
elif cardinality not in {"0/1", "1"} or order != "none":
|
|
raise SchemaError(f"{where}: scalar/enum/record fields require 0/1 or 1 and none")
|
|
if order == "none":
|
|
if order_by:
|
|
raise SchemaError(f"{where}.order_by: none fields cannot have keys")
|
|
return
|
|
if order == "ordered":
|
|
if order_by:
|
|
raise SchemaError(f"{where}.order_by: ordered fields cannot have keys")
|
|
return
|
|
if order == "raw-utf8":
|
|
if type_spec["kind"] != "map" or order_by != ["key"]:
|
|
raise SchemaError(f"{where}: raw-utf8 requires a map and the key descriptor")
|
|
return
|
|
if order != "sorted" or type_spec["kind"] != "list" or not order_by:
|
|
raise SchemaError(f"{where}: sorted order requires a keyed list")
|
|
for index, path in enumerate(order_by):
|
|
current = type_spec["item"]
|
|
if path != "value":
|
|
for component in path.split("."):
|
|
if current["kind"] != "record" or component == "":
|
|
raise SchemaError(f"{where}.order_by[{index}]: invalid field path")
|
|
fields = {
|
|
child["name"]: child for child in self.records[current["name"]]["fields"]
|
|
}
|
|
if component not in fields:
|
|
raise SchemaError(f"{where}.order_by[{index}]: unknown field path")
|
|
current = fields[component]["type"]
|
|
|
|
def _validate_semantic_value(
|
|
self,
|
|
type_spec: dict[str, Any],
|
|
value: Any,
|
|
where: str,
|
|
allow_empty_record: bool = False,
|
|
allow_empty_scalar: bool = False,
|
|
) -> None:
|
|
kind = type_spec["kind"]
|
|
if kind == "bool":
|
|
if not isinstance(value, bool):
|
|
raise SchemaError(f"{where}: expected bool")
|
|
elif kind == "bytes":
|
|
raw = _hex(value, where)
|
|
self._validate_scalar(type_spec, value, raw, where, allow_empty_scalar)
|
|
elif kind == "string":
|
|
if not isinstance(value, str) or "\x00" in value or unicodedata.normalize("NFC", value) != value:
|
|
raise SchemaError(f"{where}: expected NFC string without NUL")
|
|
self._validate_scalar(type_spec, value, value, where, allow_empty_scalar)
|
|
elif kind == "uint":
|
|
number = _uint(value, where)
|
|
self._validate_scalar(type_spec, value, number, where, allow_empty_scalar)
|
|
elif kind == "enum":
|
|
if not isinstance(value, str) or value not in {
|
|
row["name"] for row in self.enums[type_spec["name"]]["values"]
|
|
}:
|
|
raise SchemaError(f"{where}: unknown enum value")
|
|
elif kind == "record":
|
|
if not isinstance(value, dict):
|
|
raise SchemaError(f"{where}: expected record object")
|
|
if not value and allow_empty_record:
|
|
return
|
|
record = self.records[type_spec["name"]]
|
|
fields = {field["name"]: field for field in record["fields"]}
|
|
if set(value) != set(fields):
|
|
raise SchemaError(f"{where}: record default must name every encoded field")
|
|
for name, field in fields.items():
|
|
self._validate_semantic_value(
|
|
field["type"],
|
|
value[name],
|
|
where + "." + name,
|
|
allow_empty_record=field["cardinality"] == "0/1",
|
|
allow_empty_scalar=field["cardinality"] == "0/1",
|
|
)
|
|
elif kind == "list":
|
|
for index, child in enumerate(_list(value, where)):
|
|
self._validate_semantic_value(type_spec["item"], child, f"{where}[{index}]", True)
|
|
elif kind == "map":
|
|
if not isinstance(value, dict):
|
|
raise SchemaError(f"{where}: expected map object")
|
|
for key, child in value.items():
|
|
if not isinstance(key, str) or "\x00" in key or unicodedata.normalize("NFC", key) != key:
|
|
raise SchemaError(f"{where}: invalid map key")
|
|
self._validate_semantic_value(type_spec["value"], child, where + "." + key, True)
|
|
|
|
def _validate_scalar(
|
|
self,
|
|
type_spec: dict[str, Any],
|
|
original: Any,
|
|
normalized: Any,
|
|
where: str,
|
|
allow_empty: bool = False,
|
|
) -> None:
|
|
scalar_name = type_spec.get("scalar")
|
|
if not scalar_name:
|
|
return
|
|
scalar = self.scalars[scalar_name]
|
|
if allow_empty and _is_empty(original):
|
|
return
|
|
if "length" in scalar and len(normalized) != scalar["length"]:
|
|
raise SchemaError(f"{where}: {scalar_name} must have length {scalar['length']}")
|
|
if "maximum" in scalar and normalized > int(scalar["maximum"]):
|
|
raise SchemaError(f"{where}: {scalar_name} exceeds maximum")
|
|
if "values" in scalar and original not in scalar["values"]:
|
|
raise SchemaError(f"{where}: {scalar_name} is not a closed value")
|
|
if scalar.get("nonempty") and not original:
|
|
raise SchemaError(f"{where}: {scalar_name} must be nonempty")
|
|
if "path_class" in scalar:
|
|
self._validate_path(original, self.paths[scalar["path_class"]], where)
|
|
|
|
def _validate_path(self, value: str, spec: dict[str, Any], where: str) -> None:
|
|
if not isinstance(value, str) or "\\" in value or "\x00" in value:
|
|
raise SchemaError(f"{where}: invalid logical path")
|
|
if spec["ascii"] and not value.isascii():
|
|
raise SchemaError(f"{where}: path must be ASCII")
|
|
if value == "." and spec["allow_complete_dot"]:
|
|
return
|
|
absolute = value.startswith("/")
|
|
if spec["absolute"] is True and not absolute:
|
|
raise SchemaError(f"{where}: path must be absolute")
|
|
if spec["absolute"] is False and absolute:
|
|
raise SchemaError(f"{where}: path must be relative")
|
|
parts = value[1:].split("/") if absolute else value.split("/")
|
|
if spec["shape"] == "single-segment" and len(parts) != 1:
|
|
raise SchemaError(f"{where}: expected one path segment")
|
|
if any(not part or part == "." for part in parts):
|
|
raise SchemaError(f"{where}: empty or dot segment")
|
|
if not spec["allow_dotdot_segments"] and ".." in parts:
|
|
raise SchemaError(f"{where}: dot-dot segment")
|
|
if "roots" in spec and not any(value == root or value.startswith(root + "/") for root in spec["roots"]):
|
|
raise SchemaError(f"{where}: path outside virtual roots")
|
|
|
|
def _validate_digests(self) -> None:
|
|
data = self.digests_data
|
|
_keys(data, {"format", "version", "domains"}, "$digests")
|
|
if data["format"] != "wwar-v1-digest-schema" or data["version"] != 1:
|
|
raise SchemaError("$digests: unsupported format/version")
|
|
names: set[str] = set()
|
|
separators: set[bytes] = set()
|
|
for index, domain in enumerate(_list(data["domains"], "$digests.domains")):
|
|
where = f"$digests.domains[{index}]"
|
|
_allowed_keys(
|
|
domain,
|
|
{"name", "algorithm", "separator_utf8_hex", "input", "entry_encoding", "formula", "result"},
|
|
{"name", "algorithm", "separator_utf8_hex", "input", "formula", "result"},
|
|
where,
|
|
)
|
|
name = _name(domain["name"], where + ".name")
|
|
separator = _hex(domain["separator_utf8_hex"], where + ".separator_utf8_hex")
|
|
if name in names or separator in separators:
|
|
raise SchemaError(f"{where}: duplicate digest name or separator")
|
|
names.add(name)
|
|
separators.add(separator)
|
|
if domain["algorithm"] != "sha256" or not separator.endswith(b"\x00"):
|
|
raise SchemaError(f"{where}: unsupported algorithm or unterminated separator")
|
|
if len(separator) > 255:
|
|
raise SchemaError(f"{where}.separator_utf8_hex: separator is too long")
|
|
if separator.count(b"\x00") != 1:
|
|
raise SchemaError(f"{where}.separator_utf8_hex: expected one terminal NUL")
|
|
try:
|
|
separator_text = separator[:-1].decode("utf-8")
|
|
except UnicodeDecodeError as exc:
|
|
raise SchemaError(f"{where}.separator_utf8_hex: invalid UTF-8") from exc
|
|
if (
|
|
not separator_text
|
|
or unicodedata.normalize("NFC", separator_text) != separator_text
|
|
):
|
|
raise SchemaError(f"{where}.separator_utf8_hex: expected nonempty NFC UTF-8")
|
|
formula = tuple(_list(domain["formula"], where + ".formula"))
|
|
source_tree = "entry_encoding" in domain
|
|
if formula == FORMULA_BYTES:
|
|
_keys(domain["input"], {"kind"}, where + ".input")
|
|
if domain["input"]["kind"] != "bytes" or source_tree:
|
|
raise SchemaError(f"{where}: invalid bytes formula input")
|
|
elif formula == FORMULA_WWAR:
|
|
_keys(domain["input"], {"record", "encoding"}, where + ".input")
|
|
if domain["input"]["record"] not in self.records or domain["input"]["encoding"] != "WWAR" or source_tree:
|
|
raise SchemaError(f"{where}: invalid WWAR formula input")
|
|
elif formula == FORMULA_RECORD:
|
|
_keys(domain["input"], {"record", "encoding"}, where + ".input")
|
|
if domain["input"] != {"record": "top-level record", "encoding": "WWAR"} or source_tree:
|
|
raise SchemaError(f"{where}: invalid record formula input")
|
|
elif formula == FORMULA_SOURCE_TREE:
|
|
_keys(domain["input"], {"record", "encoding"}, where + ".input")
|
|
if domain["input"]["record"] not in self.records or domain["input"]["encoding"] != "special-source-tree-v1" or not source_tree:
|
|
raise SchemaError(f"{where}: invalid source-tree formula input")
|
|
self._validate_source_tree_shape(domain["entry_encoding"], where + ".entry_encoding")
|
|
else:
|
|
raise SchemaError(f"{where}.formula: unknown finite formula")
|
|
self._validate_digest_result(domain, where + ".result")
|
|
self.domains = {row["name"]: row for row in data["domains"]}
|
|
|
|
def _validate_source_tree_shape(self, value: Any, where: str) -> None:
|
|
_keys(value, {"path", "type", "executable", "content", "order"}, where)
|
|
if value["path"] != ["u64be-byte-length", "raw-nfc-utf8"]:
|
|
raise SchemaError(f"{where}.path: unsupported")
|
|
if value["type"] != {"dir": 1, "file": 2}:
|
|
raise SchemaError(f"{where}.type: unsupported")
|
|
if value["executable"] != {"false": 0, "true": 1, "directory": 0}:
|
|
raise SchemaError(f"{where}.executable: unsupported")
|
|
if value["content"] != {"dir": ["u64be", 0], "file": ["u64be", 32, "raw-sha256"]}:
|
|
raise SchemaError(f"{where}.content: unsupported")
|
|
if value["order"] != "raw-utf8-path-bytes":
|
|
raise SchemaError(f"{where}.order: unsupported")
|
|
|
|
def _validate_digest_result(self, domain: dict[str, Any], where: str) -> None:
|
|
value = domain["result"]
|
|
if not isinstance(value, dict):
|
|
raise SchemaError(f"{where}: expected object")
|
|
if "kind" in value:
|
|
if value.get("kind") != "bytes":
|
|
raise SchemaError(f"{where}.kind: digest results are bytes")
|
|
if set(value) == {"kind", "length"}:
|
|
if value["length"] != 32:
|
|
raise SchemaError(f"{where}.length: SHA-256 is 32 bytes")
|
|
return
|
|
if set(value) == {"kind", "scalar"}:
|
|
scalar = self.scalars.get(value["scalar"])
|
|
if scalar is None or scalar["base"] != "bytes" or scalar.get("length") != 32:
|
|
raise SchemaError(f"{where}.scalar: expected a 32-byte scalar")
|
|
return
|
|
raise SchemaError(f"{where}: invalid byte-result shape")
|
|
_keys(value, {"record", "domain", "record-kind", "record-schema"}, where)
|
|
if value["record"] not in self.records:
|
|
raise SchemaError(f"{where}.record: unknown result record")
|
|
if value["domain"] != domain["name"]:
|
|
raise SchemaError(f"{where}.domain: result domain differs from digest domain")
|
|
union = self.unions.get(value["record"])
|
|
if union is None or value["domain"] not in {
|
|
case["value"] for case in union["cases"]
|
|
}:
|
|
raise SchemaError(f"{where}.record: result is not the matching typed union")
|
|
if value["domain"] == "blob":
|
|
_uint(value["record-kind"], where + ".record-kind")
|
|
_uint(value["record-schema"], where + ".record-schema")
|
|
if (value["record-kind"], value["record-schema"]) != (0, 0):
|
|
raise SchemaError(f"{where}: blob identity must use kind/schema zero")
|
|
elif value["domain"] == "record":
|
|
_uint(value["record-schema"], where + ".record-schema")
|
|
if value["record-kind"] != "1..22" or value["record-schema"] != 1:
|
|
raise SchemaError(f"{where}: record identity range differs")
|
|
else:
|
|
raise SchemaError(f"{where}.domain: invalid typed digest domain")
|
|
|
|
def generated_table(self) -> dict[str, Any]:
|
|
"""Compile descriptive schemas into compact positional codec lookup tables."""
|
|
return {
|
|
"format": TABLE_FORMAT,
|
|
"version": 1,
|
|
"schema_digest": self.schema_digest,
|
|
"schema_files": self.file_hashes,
|
|
"wire": {
|
|
"magic": self.wire["magic_hex"],
|
|
"schema": self.wire["schema_version"],
|
|
"types": [
|
|
[row["code"], row["name"]] for row in self.wire["wire_types"]
|
|
],
|
|
"limits": [
|
|
self.wire["limits"]["container_members_max"],
|
|
self.wire["limits"]["nesting_depth_max"],
|
|
self.wire["limits"]["string_or_bytes_length_max"],
|
|
self.wire["limits"]["uint_max"],
|
|
self.wire["limits"]["wwar_payload_length_max"],
|
|
],
|
|
"unicode": [
|
|
self.wire["unicode"]["version"],
|
|
self.wire["unicode"]["form"],
|
|
self.wire["unicode"]["nul_forbidden"],
|
|
self.wire["unicode"]["scalar_values_only"],
|
|
],
|
|
},
|
|
"scalars": {
|
|
row["name"]: [
|
|
row["base"],
|
|
row.get("length"),
|
|
row.get("maximum"),
|
|
row.get("values"),
|
|
row.get("path_class"),
|
|
row.get("nonempty"),
|
|
]
|
|
for row in self.records_data["scalar_types"]
|
|
},
|
|
"paths": {
|
|
row["name"]: [
|
|
row["shape"],
|
|
row["absolute"],
|
|
row["allow_complete_dot"],
|
|
row["allow_dotdot_segments"],
|
|
row["ascii"],
|
|
row.get("roots"),
|
|
]
|
|
for row in self.records_data["path_classes"]
|
|
},
|
|
"enums": {
|
|
row["name"]: [[item["value"], item["name"]] for item in row["values"]]
|
|
for row in self.records_data["enums"]
|
|
},
|
|
"records": {
|
|
row["name"]: [
|
|
row["top_level_kind"],
|
|
[
|
|
[
|
|
field["tag"],
|
|
field["name"],
|
|
_compiled_type(field["type"]),
|
|
field["cardinality"],
|
|
field["encoded_default"],
|
|
field["order"],
|
|
field["order_by"],
|
|
]
|
|
for field in row["fields"]
|
|
],
|
|
]
|
|
for row in self.records_data["records"]
|
|
},
|
|
"unions": {
|
|
row["record"]: [
|
|
row["discriminator"],
|
|
{
|
|
case["value"]: [
|
|
case["required"],
|
|
case["allowed"],
|
|
case["nonempty"],
|
|
case["empty"],
|
|
]
|
|
for case in row["cases"]
|
|
},
|
|
]
|
|
for row in self.records_data["unions"]
|
|
},
|
|
"kinds": {
|
|
str(row["kind"]): [row["name"], row["schema"]]
|
|
for row in self.records_data["record_kinds"]
|
|
},
|
|
"artifacts": {
|
|
row["artifact_kind"]: [
|
|
row["domain"],
|
|
row["record_kind"],
|
|
row["record_schema"],
|
|
]
|
|
for row in self.records_data["artifact_digest_mapping"]
|
|
},
|
|
"wrappers": {
|
|
row["name"]: [row["magic_hex"], row["body_record"]]
|
|
for row in self.records_data["wrappers"]
|
|
},
|
|
"digests": {
|
|
row["name"]: [
|
|
row["algorithm"],
|
|
row["separator_utf8_hex"],
|
|
row["input"],
|
|
row["formula"],
|
|
row.get("entry_encoding"),
|
|
]
|
|
for row in self.digests_data["domains"]
|
|
},
|
|
}
|
|
|
|
def _materialize_record(self, name: str, value: Any, path: str) -> dict[str, Any]:
|
|
if name not in self.records or not isinstance(value, dict):
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path, "expected declared record object")
|
|
record = self.records[name]
|
|
fields = {field["name"]: field for field in record["fields"]}
|
|
unknown = sorted(set(value) - set(fields))
|
|
if unknown:
|
|
raise ProtocolError("UNKNOWN_FIELD", path + "." + unknown[0])
|
|
result: dict[str, Any] = {}
|
|
for field in record["fields"]:
|
|
field_name = field["name"]
|
|
if field_name in value:
|
|
result[field_name] = copy.deepcopy(value[field_name])
|
|
elif field["encoded_default"] is not None:
|
|
result[field_name] = copy.deepcopy(field["encoded_default"])
|
|
else:
|
|
raise ProtocolError("MISSING_FIELD", path + "." + field_name)
|
|
union = self.unions.get(name)
|
|
if union:
|
|
self._validate_union(union, result, path)
|
|
return result
|
|
|
|
def _validate_union(self, union: dict[str, Any], value: dict[str, Any], path: str) -> None:
|
|
discriminator = union["discriminator"]
|
|
selected = value[discriminator]
|
|
case = next((row for row in union["cases"] if row["value"] == selected), None)
|
|
if case is None:
|
|
raise ProtocolError("INVALID_UNION", path + "." + discriminator)
|
|
for field_name in case["nonempty"]:
|
|
if _is_empty(value[field_name]):
|
|
raise ProtocolError("INVALID_UNION", path + "." + field_name)
|
|
for field_name in case["empty"]:
|
|
if not _is_empty(value[field_name]):
|
|
raise ProtocolError("INVALID_UNION", path + "." + field_name)
|
|
|
|
def _order_component(
|
|
self, value: Any, type_spec: dict[str, Any], field_path: str, path: str
|
|
) -> tuple[Any, dict[str, Any]]:
|
|
if field_path == "value":
|
|
return value, type_spec
|
|
current_value = value
|
|
current_type = type_spec
|
|
for component in field_path.split("."):
|
|
if current_type["kind"] != "record":
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path, "order path is not a record")
|
|
current_value = self._materialize_record(
|
|
current_type["name"], current_value, path
|
|
)
|
|
fields = {
|
|
field["name"]: field for field in self.records[current_type["name"]]["fields"]
|
|
}
|
|
field = fields[component]
|
|
current_value = current_value[component]
|
|
current_type = field["type"]
|
|
return current_value, current_type
|
|
|
|
def _canonical_order_key(
|
|
self, value: Any, type_spec: dict[str, Any], path: str
|
|
) -> tuple[Any, ...]:
|
|
kind = type_spec["kind"]
|
|
if kind == "string":
|
|
return (0, value.encode("utf-8"))
|
|
if kind == "bytes":
|
|
return (1, bytes.fromhex(value))
|
|
if kind == "uint":
|
|
return (2, int(value))
|
|
if kind == "enum":
|
|
member = next(
|
|
row for row in self.enums[type_spec["name"]]["values"] if row["name"] == value
|
|
)
|
|
return (2, member["value"])
|
|
if kind == "bool":
|
|
return (3, int(value))
|
|
if kind == "list":
|
|
return (
|
|
4,
|
|
tuple(
|
|
self._canonical_order_key(child, type_spec["item"], f"{path}[{index}]")
|
|
for index, child in enumerate(value)
|
|
),
|
|
)
|
|
if kind == "map":
|
|
return (
|
|
5,
|
|
tuple(
|
|
(
|
|
key.encode("utf-8"),
|
|
self._canonical_order_key(value[key], type_spec["value"], path + "." + key),
|
|
)
|
|
for key in sorted(value, key=lambda item: item.encode("utf-8"))
|
|
),
|
|
)
|
|
if kind == "record":
|
|
return (6, self.encode_record(type_spec["name"], value))
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path, "unknown order-key type")
|
|
|
|
def _validate_field_order(
|
|
self, field: dict[str, Any], value: Any, path: str
|
|
) -> None:
|
|
if field["order"] != "sorted":
|
|
return
|
|
keys = []
|
|
for index, item in enumerate(value):
|
|
components = []
|
|
for field_path in field["order_by"]:
|
|
component, component_type = self._order_component(
|
|
item, field["type"]["item"], field_path, f"{path}[{index}]"
|
|
)
|
|
components.append(
|
|
self._canonical_order_key(
|
|
component, component_type, f"{path}[{index}].{field_path}"
|
|
)
|
|
)
|
|
keys.append(tuple(components))
|
|
for index, (left, right) in enumerate(zip(keys, keys[1:]), 1):
|
|
if right <= left:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", f"{path}[{index}]", "list order")
|
|
|
|
def conformance_sample(self, name: str, union_case: str | None = None) -> dict[str, Any]:
|
|
"""Construct one generic valid value for codec coverage, never build policy."""
|
|
return self._sample_record(name, (), union_case)
|
|
|
|
def _sample_type(self, type_spec: dict[str, Any], stack: tuple[str, ...]) -> Any:
|
|
kind = type_spec["kind"]
|
|
scalar = self.scalars.get(type_spec.get("scalar", ""), {})
|
|
if kind == "bool":
|
|
return False
|
|
if kind == "bytes":
|
|
return "00" * scalar.get("length", 0)
|
|
if kind == "string":
|
|
if scalar.get("values"):
|
|
return scalar["values"][0]
|
|
if scalar.get("path_class"):
|
|
path = self.paths[scalar["path_class"]]
|
|
if path["allow_complete_dot"]:
|
|
return "."
|
|
if path["absolute"] is True:
|
|
return path.get("roots", ["/x"])[0]
|
|
return "x"
|
|
return "x" if scalar.get("nonempty") else ""
|
|
if kind == "uint":
|
|
return scalar.get("values", [0])[0]
|
|
if kind == "enum":
|
|
return self.enums[type_spec["name"]]["values"][0]["name"]
|
|
if kind == "record":
|
|
return self._sample_record(type_spec["name"], stack)
|
|
if kind == "list":
|
|
return []
|
|
if kind == "map":
|
|
return {}
|
|
raise SchemaError(f"unknown sample type {kind!r}")
|
|
|
|
def _nonempty_sample(self, type_spec: dict[str, Any], stack: tuple[str, ...]) -> Any:
|
|
kind = type_spec["kind"]
|
|
if kind == "bool":
|
|
return True
|
|
if kind == "bytes":
|
|
scalar = self.scalars.get(type_spec.get("scalar", ""), {})
|
|
return "00" * scalar.get("length", 1)
|
|
if kind == "string":
|
|
return self._sample_type(type_spec, stack) or "x"
|
|
if kind == "uint":
|
|
return 1
|
|
if kind == "enum":
|
|
return self._sample_type(type_spec, stack)
|
|
if kind == "record":
|
|
name = type_spec["name"]
|
|
if name in stack:
|
|
union = self.unions.get(name)
|
|
if union is None:
|
|
raise SchemaError(f"required recursive sample for {name}")
|
|
record = self.records[name]
|
|
fields = {field["name"]: field for field in record["fields"]}
|
|
|
|
def recursive(spec: dict[str, Any]) -> bool:
|
|
if spec["kind"] == "record":
|
|
return spec["name"] == name
|
|
if spec["kind"] == "list":
|
|
return recursive(spec["item"])
|
|
if spec["kind"] == "map":
|
|
return recursive(spec["value"])
|
|
return False
|
|
|
|
case = next(
|
|
(
|
|
row
|
|
for row in union["cases"]
|
|
if not any(recursive(fields[field]["type"]) for field in row["nonempty"])
|
|
),
|
|
None,
|
|
)
|
|
if case is None:
|
|
raise SchemaError(f"no finite recursive sample for {name}")
|
|
return self._sample_record(name, (), case["value"])
|
|
return self._sample_record(name, stack)
|
|
if kind == "list":
|
|
return [self._nonempty_sample(type_spec["item"], stack)]
|
|
if kind == "map":
|
|
return {"x": self._nonempty_sample(type_spec["value"], stack)}
|
|
raise SchemaError(f"unknown nonempty sample type {kind!r}")
|
|
|
|
def _sample_record(
|
|
self, name: str, stack: tuple[str, ...], union_case: str | None = None
|
|
) -> dict[str, Any]:
|
|
if name in stack:
|
|
raise SchemaError(f"required record recursion: {' -> '.join(stack + (name,))}")
|
|
record = self.records.get(name)
|
|
if record is None:
|
|
raise SchemaError(f"unknown sample record {name!r}")
|
|
value: dict[str, Any] = {}
|
|
for field in record["fields"]:
|
|
default = field["encoded_default"]
|
|
value[field["name"]] = (
|
|
copy.deepcopy(default)
|
|
if default is not None
|
|
else self._sample_type(field["type"], stack + (name,))
|
|
)
|
|
union = self.unions.get(name)
|
|
if union:
|
|
selected = union_case or union["cases"][0]["value"]
|
|
case = next((row for row in union["cases"] if row["value"] == selected), None)
|
|
if case is None:
|
|
raise SchemaError(f"unknown sample union case {name}.{selected}")
|
|
value[union["discriminator"]] = selected
|
|
fields = {field["name"]: field for field in record["fields"]}
|
|
for field_name in case["empty"]:
|
|
default = fields[field_name]["encoded_default"]
|
|
if default is None:
|
|
default = {
|
|
"bool": False,
|
|
"bytes": "",
|
|
"string": "",
|
|
"uint": 0,
|
|
"list": [],
|
|
"map": {},
|
|
"record": {},
|
|
}[fields[field_name]["type"]["kind"]]
|
|
value[field_name] = copy.deepcopy(default)
|
|
for field_name in case["nonempty"]:
|
|
value[field_name] = self._nonempty_sample(
|
|
fields[field_name]["type"], stack + (name,)
|
|
)
|
|
return value
|
|
|
|
def _semantic_ast(
|
|
self,
|
|
type_spec: dict[str, Any],
|
|
value: Any,
|
|
path: str,
|
|
cardinality: str = "1",
|
|
) -> dict[str, Any]:
|
|
kind = type_spec["kind"]
|
|
try:
|
|
if kind == "bool":
|
|
if not isinstance(value, bool):
|
|
raise SchemaError("expected bool")
|
|
return {"bool": value}
|
|
if kind == "bytes":
|
|
raw = _hex(value, path)
|
|
self._validate_scalar(type_spec, value, raw, path, cardinality == "0/1")
|
|
return {"bytes": value}
|
|
if kind == "string":
|
|
if not isinstance(value, str) or "\x00" in value or unicodedata.normalize("NFC", value) != value:
|
|
raise SchemaError("expected NFC string without NUL")
|
|
self._validate_scalar(type_spec, value, value, path, cardinality == "0/1")
|
|
return {"string": value}
|
|
if kind == "uint":
|
|
number = _uint(value, path)
|
|
self._validate_scalar(type_spec, value, number, path, cardinality == "0/1")
|
|
return {"uint": str(number)}
|
|
if kind == "enum":
|
|
enum = self.enums[type_spec["name"]]
|
|
member = next((row for row in enum["values"] if row["name"] == value), None)
|
|
if member is None:
|
|
raise ProtocolError("UNKNOWN_ENUM", path)
|
|
return {"uint": str(member["value"])}
|
|
if kind == "record":
|
|
if value == {} and cardinality == "0/1":
|
|
return {"record": []}
|
|
return self._record_ast(type_spec["name"], value, path)
|
|
if kind == "list":
|
|
if not isinstance(value, list):
|
|
raise SchemaError("expected list")
|
|
return {
|
|
"list": [
|
|
self._semantic_ast(type_spec["item"], child, f"{path}[{index}]")
|
|
for index, child in enumerate(value)
|
|
]
|
|
}
|
|
if kind == "map":
|
|
if not isinstance(value, dict):
|
|
raise SchemaError("expected map object")
|
|
pairs = []
|
|
for key in sorted(value, key=lambda item: item.encode("utf-8")):
|
|
_validate_string(key, path + ".key")
|
|
pairs.append([key, self._semantic_ast(type_spec["value"], value[key], path + "." + key)])
|
|
return {"map": pairs}
|
|
except SchemaError as exc:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path, str(exc)) from exc
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path, "unknown field type")
|
|
|
|
def _record_ast(self, name: str, value: Any, path: str = "$") -> dict[str, Any]:
|
|
materialized = self._materialize_record(name, value, path)
|
|
fields = []
|
|
for field in self.records[name]["fields"]:
|
|
field_path = path + "." + field["name"]
|
|
self._validate_field_order(field, materialized[field["name"]], field_path)
|
|
fields.append(
|
|
[
|
|
field["tag"],
|
|
self._semantic_ast(
|
|
field["type"], materialized[field["name"]], field_path, field["cardinality"]
|
|
),
|
|
]
|
|
)
|
|
return {"record": fields}
|
|
|
|
def encode_record(self, name: str, value: Any) -> bytes:
|
|
return encode_envelope(self._record_ast(name, value), self.wire)
|
|
|
|
def decode_record(self, name: str, encoded: bytes) -> dict[str, Any]:
|
|
ast = decode_envelope(encoded, self.wire)
|
|
return self._decode_record_ast(name, ast, "$")
|
|
|
|
def _decode_record_ast(self, name: str, ast: Any, path: str) -> dict[str, Any]:
|
|
if not isinstance(ast, dict) or set(ast) != {"record"}:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path, "expected record wire type")
|
|
pairs = ast["record"]
|
|
fields = self.records[name]["fields"]
|
|
known = {field["tag"]: field for field in fields}
|
|
actual = {tag for tag, _ in pairs}
|
|
for tag, _ in pairs:
|
|
if tag not in known:
|
|
raise ProtocolError("UNKNOWN_FIELD", path + f".{tag}")
|
|
for field in fields:
|
|
if field["tag"] not in actual:
|
|
raise ProtocolError("MISSING_FIELD", path + "." + field["name"])
|
|
result: dict[str, Any] = {}
|
|
for tag, child in pairs:
|
|
field = known[tag]
|
|
result[field["name"]] = self._decode_semantic_ast(
|
|
field["type"], child, path + "." + field["name"], field["cardinality"]
|
|
)
|
|
for field in fields:
|
|
self._validate_field_order(field, result[field["name"]], path + "." + field["name"])
|
|
union = self.unions.get(name)
|
|
if union:
|
|
self._validate_union(union, result, path)
|
|
return result
|
|
|
|
def _decode_semantic_ast(
|
|
self, type_spec: dict[str, Any], ast: Any, path: str, cardinality: str = "1"
|
|
) -> Any:
|
|
kind = type_spec["kind"]
|
|
expected = "uint" if kind == "enum" else kind
|
|
if not isinstance(ast, dict) or set(ast) != {expected}:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path, f"expected {expected} wire type")
|
|
value = ast[expected]
|
|
try:
|
|
if kind == "bool":
|
|
return value
|
|
if kind == "bytes":
|
|
raw = _hex(value, path)
|
|
self._validate_scalar(type_spec, value, raw, path, cardinality == "0/1")
|
|
return value
|
|
if kind == "string":
|
|
self._validate_scalar(type_spec, value, value, path, cardinality == "0/1")
|
|
return value
|
|
if kind == "uint":
|
|
number = int(value)
|
|
self._validate_scalar(type_spec, number, number, path, cardinality == "0/1")
|
|
return number
|
|
if kind == "enum":
|
|
number = int(value)
|
|
member = next(
|
|
(row for row in self.enums[type_spec["name"]]["values"] if row["value"] == number),
|
|
None,
|
|
)
|
|
if member is None:
|
|
raise ProtocolError("UNKNOWN_ENUM", path)
|
|
return member["name"]
|
|
if kind == "record":
|
|
if value == [] and cardinality == "0/1":
|
|
return {}
|
|
return self._decode_record_ast(type_spec["name"], ast, path)
|
|
if kind == "list":
|
|
return [
|
|
self._decode_semantic_ast(type_spec["item"], child, f"{path}[{index}]")
|
|
for index, child in enumerate(value)
|
|
]
|
|
if kind == "map":
|
|
return {
|
|
key: self._decode_semantic_ast(type_spec["value"], child, path + "." + key)
|
|
for key, child in value
|
|
}
|
|
except SchemaError as exc:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path, str(exc)) from exc
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path, "unknown field type")
|
|
|
|
def record_id(self, name: str, kind: int, schema: int, value: Any) -> bytes:
|
|
record = self.records.get(name)
|
|
kind_row = self.kinds.get(kind)
|
|
if (
|
|
record is None
|
|
or kind_row is None
|
|
or kind_row["name"] != name
|
|
or kind_row["schema"] != schema
|
|
or record["top_level_kind"] != kind
|
|
):
|
|
raise ProtocolError("KIND_SUBSTITUTION", "$.kind")
|
|
encoded = self.encode_record(name, value)
|
|
domain = self._domain_by_formula(FORMULA_RECORD)
|
|
return hashlib.sha256(
|
|
bytes.fromhex(domain["separator_utf8_hex"])
|
|
+ struct.pack(">I", kind)
|
|
+ struct.pack(">I", schema)
|
|
+ struct.pack(">Q", len(encoded))
|
|
+ encoded
|
|
).digest()
|
|
|
|
def digest(self, domain_name: str, value: Any, record_name: str | None = None, kind: int | None = None, schema: int = 1) -> bytes:
|
|
domain = self.domains.get(domain_name)
|
|
if domain is None:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", "$.domain", "unknown digest domain")
|
|
formula = tuple(domain["formula"])
|
|
separator = bytes.fromhex(domain["separator_utf8_hex"])
|
|
if formula == FORMULA_BYTES:
|
|
if not isinstance(value, bytes):
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", "$", "bytes digest needs bytes")
|
|
preimage = separator + struct.pack(">Q", len(value)) + value
|
|
elif formula == FORMULA_WWAR:
|
|
expected_record = domain["input"]["record"]
|
|
if record_name is not None and record_name != expected_record:
|
|
raise ProtocolError("KIND_SUBSTITUTION", "$.record")
|
|
encoded = self.encode_record(expected_record, value)
|
|
preimage = separator + struct.pack(">Q", len(encoded)) + encoded
|
|
elif formula == FORMULA_RECORD:
|
|
if record_name is None or kind is None:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", "$", "record digest needs record and kind")
|
|
return self.record_id(record_name, kind, schema, value)
|
|
elif formula == FORMULA_SOURCE_TREE:
|
|
encoded_entries = self._encode_source_tree(value)
|
|
preimage = separator + encoded_entries
|
|
else:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", "$.domain", "unsupported digest formula")
|
|
return hashlib.sha256(preimage).digest()
|
|
|
|
def _domain_by_formula(self, formula: tuple[str, ...]) -> dict[str, Any]:
|
|
matches = [row for row in self.domains.values() if tuple(row["formula"]) == formula]
|
|
if len(matches) != 1:
|
|
raise SchemaError(f"expected one domain for formula {formula}")
|
|
return matches[0]
|
|
|
|
def _encode_source_tree(self, entries: Any) -> bytes:
|
|
if not isinstance(entries, list):
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", "$", "source tree needs an entry list")
|
|
out = bytearray()
|
|
previous: bytes | None = None
|
|
for index, entry in enumerate(entries):
|
|
path = f"$[{index}]"
|
|
if not isinstance(entry, dict) or set(entry) != {"path", "type", "executable", "content"}:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path, "invalid source-tree entry")
|
|
name = entry["path"]
|
|
try:
|
|
_validate_string(name, path + ".path")
|
|
self._validate_path(name, {
|
|
"ascii": False,
|
|
"allow_complete_dot": False,
|
|
"shape": "path",
|
|
"absolute": False,
|
|
"allow_dotdot_segments": False,
|
|
}, path + ".path")
|
|
except SchemaError as exc:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path + ".path", str(exc)) from exc
|
|
raw_path = name.encode("utf-8")
|
|
if previous is not None and raw_path <= previous:
|
|
code = "DUPLICATE_MAP_KEY" if raw_path == previous else "MAP_KEY_ORDER"
|
|
raise ProtocolError(code, path + ".path")
|
|
previous = raw_path
|
|
out += struct.pack(">Q", len(raw_path)) + raw_path
|
|
if entry["type"] == "dir":
|
|
if entry["executable"] is not False or entry["content"] != "":
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path)
|
|
out += b"\x01\x00" + struct.pack(">Q", 0)
|
|
elif entry["type"] == "file":
|
|
if not isinstance(entry["executable"], bool):
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path + ".executable")
|
|
try:
|
|
content = _hex(entry["content"], path + ".content", 32)
|
|
except SchemaError as exc:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path + ".content", str(exc)) from exc
|
|
out += b"\x02" + bytes([int(entry["executable"])]) + struct.pack(">Q", 32) + content
|
|
else:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path + ".type")
|
|
return bytes(out)
|
|
|
|
def encode_wrapper(self, wrapper_name: str, value: Any) -> bytes:
|
|
wrapper = self.wrappers.get(wrapper_name)
|
|
if wrapper is None:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", "$.wrapper")
|
|
return bytes.fromhex(wrapper["magic_hex"]) + self.encode_record(wrapper["body_record"], value)
|
|
|
|
def decode_wrapper(self, wrapper_name: str, encoded: bytes) -> dict[str, Any]:
|
|
wrapper = self.wrappers.get(wrapper_name)
|
|
if wrapper is None:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", "$.wrapper")
|
|
magic = bytes.fromhex(wrapper["magic_hex"])
|
|
if len(encoded) < len(magic):
|
|
raise ProtocolError("TRUNCATED")
|
|
if encoded[: len(magic)] != magic:
|
|
raise ProtocolError("BAD_MAGIC")
|
|
return self.decode_record(wrapper["body_record"], encoded[len(magic) :])
|
|
|
|
|
|
def _is_empty(value: Any) -> bool:
|
|
return value is False or value == 0 or value == "" or value == [] or value == {}
|
|
|
|
|
|
def _validate_string(value: Any, where: str) -> None:
|
|
if not isinstance(value, str):
|
|
raise SchemaError(f"{where}: expected string")
|
|
if "\x00" in value:
|
|
raise SchemaError(f"{where}: NUL is forbidden")
|
|
if unicodedata.normalize("NFC", value) != value:
|
|
raise SchemaError(f"{where}: string is not NFC")
|
|
|
|
|
|
def schema_bundle_digest(paths: dict[str, pathlib.Path]) -> str:
|
|
hasher = hashlib.sha256()
|
|
hasher.update(SCHEMA_DOMAIN)
|
|
for name in SCHEMA_FILES:
|
|
raw_name = name.encode("utf-8")
|
|
raw = paths[name].read_bytes()
|
|
hasher.update(struct.pack(">Q", len(raw_name)))
|
|
hasher.update(raw_name)
|
|
hasher.update(struct.pack(">Q", len(raw)))
|
|
hasher.update(raw)
|
|
return hasher.hexdigest()
|
|
|
|
|
|
def write_schema_manifest(schema_dir: pathlib.Path) -> None:
|
|
lines = []
|
|
for name in SCHEMA_FILES:
|
|
digest = hashlib.sha256((schema_dir / name).read_bytes()).hexdigest()
|
|
lines.append(f"{digest} {name}\n")
|
|
(schema_dir / SCHEMA_MANIFEST).write_text("".join(lines), encoding="ascii")
|
|
|
|
|
|
def check_schema_manifest(schema_dir: pathlib.Path, expected: dict[str, str] | None = None) -> None:
|
|
path = schema_dir / SCHEMA_MANIFEST
|
|
try:
|
|
lines = path.read_text(encoding="ascii").splitlines()
|
|
except (OSError, UnicodeError) as exc:
|
|
raise SchemaError(f"cannot read schema manifest: {exc}") from exc
|
|
if len(lines) != len(SCHEMA_FILES):
|
|
raise SchemaError("schema manifest must contain exactly three lines")
|
|
parsed: dict[str, str] = {}
|
|
for index, line in enumerate(lines):
|
|
parts = line.split(" ")
|
|
if len(parts) != 2 or len(parts[0]) != 64 or parts[0].lower() != parts[0]:
|
|
raise SchemaError(f"schema manifest line {index + 1} is not canonical")
|
|
_hex(parts[0], f"schema manifest line {index + 1}", 32)
|
|
if parts[1] in parsed:
|
|
raise SchemaError("schema manifest has duplicate file")
|
|
parsed[parts[1]] = parts[0]
|
|
actual = expected or {
|
|
name: hashlib.sha256((schema_dir / name).read_bytes()).hexdigest() for name in SCHEMA_FILES
|
|
}
|
|
if list(parsed) != list(SCHEMA_FILES) or parsed != actual:
|
|
raise SchemaError("schema manifest does not match exact schema bytes")
|
|
|
|
|
|
def generate_table(schema_dir: pathlib.Path, output: pathlib.Path) -> bytes:
|
|
bundle = SchemaBundle.from_dir(schema_dir)
|
|
raw = canonical_json(bundle.generated_table())
|
|
output.parent.mkdir(parents=True, exist_ok=True)
|
|
output.write_bytes(raw)
|
|
return raw
|
|
|
|
|
|
def expected_table(schema_dir: pathlib.Path) -> bytes:
|
|
return canonical_json(SchemaBundle.from_dir(schema_dir).generated_table())
|
|
|
|
|
|
def encode_envelope(ast: Any, wire: dict[str, Any]) -> bytes:
|
|
payload = _encode_ast(ast, wire, 1, "$")
|
|
return bytes.fromhex(wire["magic_hex"]) + struct.pack(">H", wire["schema_version"]) + payload
|
|
|
|
|
|
def _encode_ast(ast: Any, wire: dict[str, Any], depth: int, path: str) -> bytes:
|
|
if depth > wire["limits"]["nesting_depth_max"]:
|
|
raise ProtocolError("LIMIT_EXCEEDED", path)
|
|
if not isinstance(ast, dict) or len(ast) != 1:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path, "value AST needs one type member")
|
|
name, value = next(iter(ast.items()))
|
|
codes = {row["name"]: row["code"] for row in wire["wire_types"]}
|
|
if name not in codes:
|
|
raise ProtocolError("UNKNOWN_TYPE", path)
|
|
limit = wire["limits"]["string_or_bytes_length_max"]
|
|
members_limit = wire["limits"]["container_members_max"]
|
|
if name == "bool":
|
|
if not isinstance(value, bool):
|
|
raise ProtocolError("INVALID_BOOL", path)
|
|
payload = bytes([int(value)])
|
|
elif name == "bytes":
|
|
try:
|
|
payload = _hex(value, path)
|
|
except SchemaError as exc:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path, str(exc)) from exc
|
|
if len(payload) > limit:
|
|
raise ProtocolError("LIMIT_EXCEEDED", path)
|
|
elif name == "string":
|
|
try:
|
|
_validate_string(value, path)
|
|
except SchemaError as exc:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path, str(exc)) from exc
|
|
payload = value.encode("utf-8")
|
|
if len(payload) > limit:
|
|
raise ProtocolError("LIMIT_EXCEEDED", path)
|
|
elif name == "uint":
|
|
try:
|
|
number = int(value)
|
|
except (TypeError, ValueError) as exc:
|
|
raise ProtocolError("NONMINIMAL_UINT", path) from exc
|
|
if isinstance(value, bool) or number < 0 or number > int(wire["limits"]["uint_max"]):
|
|
raise ProtocolError("OVERFLOW", path)
|
|
size = max(1, (number.bit_length() + 7) // 8)
|
|
payload = number.to_bytes(size, "big")
|
|
elif name == "list":
|
|
if not isinstance(value, list):
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", path)
|
|
if len(value) > members_limit:
|
|
raise ProtocolError("LIMIT_EXCEEDED", path)
|
|
body = bytearray(struct.pack(">I", len(value)))
|
|
for index, child in enumerate(value):
|
|
encoded = _encode_ast(child, wire, depth + 1, f"{path}[{index}]")
|
|
body += struct.pack(">Q", len(encoded)) + encoded
|
|
payload = bytes(body)
|
|
elif name == "map":
|
|
if not isinstance(value, list) or len(value) > members_limit:
|
|
raise ProtocolError("LIMIT_EXCEEDED" if isinstance(value, list) else "CONSTRAINT_VIOLATION", path)
|
|
body = bytearray(struct.pack(">I", len(value)))
|
|
previous: bytes | None = None
|
|
for index, pair in enumerate(value):
|
|
if not isinstance(pair, list) or len(pair) != 2:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", f"{path}[{index}]")
|
|
key, child = pair
|
|
try:
|
|
_validate_string(key, f"{path}[{index}].key")
|
|
except SchemaError as exc:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", f"{path}[{index}].key", str(exc)) from exc
|
|
raw_key = key.encode("utf-8")
|
|
if len(raw_key) > limit:
|
|
raise ProtocolError("LIMIT_EXCEEDED", f"{path}[{index}].key")
|
|
if previous is not None and raw_key <= previous:
|
|
code = "DUPLICATE_MAP_KEY" if raw_key == previous else "MAP_KEY_ORDER"
|
|
raise ProtocolError(code, f"{path}[{index}].key")
|
|
previous = raw_key
|
|
encoded = _encode_ast(child, wire, depth + 1, f"{path}[{index}]")
|
|
body += struct.pack(">Q", len(raw_key)) + raw_key + struct.pack(">Q", len(encoded)) + encoded
|
|
payload = bytes(body)
|
|
else:
|
|
if not isinstance(value, list) or len(value) > members_limit:
|
|
raise ProtocolError("LIMIT_EXCEEDED" if isinstance(value, list) else "CONSTRAINT_VIOLATION", path)
|
|
body = bytearray(struct.pack(">I", len(value)))
|
|
previous = 0
|
|
for index, pair in enumerate(value):
|
|
if not isinstance(pair, list) or len(pair) != 2:
|
|
raise ProtocolError("CONSTRAINT_VIOLATION", f"{path}.{index}")
|
|
tag, child = pair
|
|
if isinstance(tag, bool) or not isinstance(tag, int) or tag <= previous or tag > 0xFFFFFFFF:
|
|
raise ProtocolError("FIELD_ORDER", f"{path}.{tag}")
|
|
previous = tag
|
|
encoded = _encode_ast(child, wire, depth + 1, f"{path}.{tag}")
|
|
body += struct.pack(">I", tag) + struct.pack(">Q", len(encoded)) + encoded
|
|
payload = bytes(body)
|
|
return bytes([codes[name]]) + struct.pack(">Q", len(payload)) + payload
|
|
|
|
|
|
def decode_envelope(encoded: bytes, wire: dict[str, Any]) -> dict[str, Any]:
|
|
magic = bytes.fromhex(wire["magic_hex"])
|
|
if len(encoded) < len(magic):
|
|
raise ProtocolError("TRUNCATED")
|
|
if encoded[: len(magic)] != magic:
|
|
raise ProtocolError("BAD_MAGIC")
|
|
if len(encoded) < len(magic) + 2:
|
|
raise ProtocolError("TRUNCATED")
|
|
version = struct.unpack_from(">H", encoded, len(magic))[0]
|
|
if version != wire["schema_version"]:
|
|
raise ProtocolError("BAD_VERSION")
|
|
ast, end = _decode_ast(encoded, len(magic) + 2, len(encoded), wire, 1, "$", False)
|
|
if end != len(encoded):
|
|
raise ProtocolError("TRAILING_BYTES")
|
|
return ast
|
|
|
|
|
|
def _need(data: bytes, offset: int, size: int, limit: int, path: str, bounded: bool = False) -> None:
|
|
if size < 0 or offset + size > limit:
|
|
raise ProtocolError("LENGTH_MISMATCH" if bounded else "TRUNCATED", path)
|
|
if offset + size > len(data):
|
|
raise ProtocolError("TRUNCATED", path)
|
|
|
|
|
|
def _decode_ast(
|
|
data: bytes,
|
|
offset: int,
|
|
limit: int,
|
|
wire: dict[str, Any],
|
|
depth: int,
|
|
path: str,
|
|
bounded: bool,
|
|
) -> tuple[dict[str, Any], int]:
|
|
if depth > wire["limits"]["nesting_depth_max"]:
|
|
raise ProtocolError("LIMIT_EXCEEDED", path)
|
|
_need(data, offset, 9, limit, path, bounded)
|
|
code = data[offset]
|
|
length = struct.unpack_from(">Q", data, offset + 1)[0]
|
|
names = {row["code"]: row["name"] for row in wire["wire_types"]}
|
|
if code not in names:
|
|
raise ProtocolError("UNKNOWN_TYPE", path)
|
|
name = names[code]
|
|
if name in {"bytes", "string"} and length > wire["limits"]["string_or_bytes_length_max"]:
|
|
raise ProtocolError("LIMIT_EXCEEDED", path)
|
|
start = offset + 9
|
|
end = start + length
|
|
if end > limit:
|
|
raise ProtocolError("LENGTH_MISMATCH" if bounded else "TRUNCATED", path)
|
|
if end > len(data):
|
|
raise ProtocolError("TRUNCATED", path)
|
|
payload = data[start:end]
|
|
if name == "bytes":
|
|
value: Any = payload.hex()
|
|
elif name == "string":
|
|
try:
|
|
value = payload.decode("utf-8")
|
|
except UnicodeDecodeError as exc:
|
|
raise ProtocolError("INVALID_UTF8", path) from exc
|
|
if "\x00" in value:
|
|
raise ProtocolError("NUL_STRING", path)
|
|
if unicodedata.normalize("NFC", value) != value:
|
|
raise ProtocolError("NON_NFC", path)
|
|
elif name == "uint":
|
|
if length > 8:
|
|
raise ProtocolError("OVERFLOW", path)
|
|
if length == 0 or (length > 1 and payload[0] == 0):
|
|
raise ProtocolError("NONMINIMAL_UINT", path)
|
|
value = str(int.from_bytes(payload, "big"))
|
|
elif name == "bool":
|
|
if length != 1:
|
|
raise ProtocolError("LENGTH_MISMATCH", path)
|
|
if payload[0] not in (0, 1):
|
|
raise ProtocolError("INVALID_BOOL", path)
|
|
value = bool(payload[0])
|
|
elif name == "list":
|
|
value = _decode_list(data, start, end, wire, depth, path)
|
|
elif name == "map":
|
|
value = _decode_map(data, start, end, wire, depth, path)
|
|
else:
|
|
value = _decode_record(data, start, end, wire, depth, path)
|
|
return {name: value}, end
|
|
|
|
|
|
def _container_count(data: bytes, start: int, end: int, wire: dict[str, Any], path: str) -> tuple[int, int]:
|
|
_need(data, start, 4, end, path)
|
|
count = struct.unpack_from(">I", data, start)[0]
|
|
if count > wire["limits"]["container_members_max"]:
|
|
raise ProtocolError("LIMIT_EXCEEDED", path)
|
|
return count, start + 4
|
|
|
|
|
|
def _decode_list(
|
|
data: bytes, start: int, end: int, wire: dict[str, Any], depth: int, path: str
|
|
) -> list[Any]:
|
|
count, offset = _container_count(data, start, end, wire, path)
|
|
result = []
|
|
for index in range(count):
|
|
child_path = f"{path}[{index}]"
|
|
_need(data, offset, 8, end, child_path)
|
|
length = struct.unpack_from(">Q", data, offset)[0]
|
|
offset += 8
|
|
child_end = offset + length
|
|
if child_end > end:
|
|
raise ProtocolError("LENGTH_MISMATCH", child_path)
|
|
child, consumed = _decode_ast(data, offset, child_end, wire, depth + 1, child_path, True)
|
|
if consumed != child_end:
|
|
raise ProtocolError("LENGTH_MISMATCH", child_path)
|
|
result.append(child)
|
|
offset = child_end
|
|
if offset != end:
|
|
raise ProtocolError("LENGTH_MISMATCH", path)
|
|
return result
|
|
|
|
|
|
def _decode_map(
|
|
data: bytes, start: int, end: int, wire: dict[str, Any], depth: int, path: str
|
|
) -> list[Any]:
|
|
count, offset = _container_count(data, start, end, wire, path)
|
|
result = []
|
|
previous: bytes | None = None
|
|
for index in range(count):
|
|
key_path = f"{path}[{index}].key"
|
|
_need(data, offset, 8, end, key_path)
|
|
key_length = struct.unpack_from(">Q", data, offset)[0]
|
|
offset += 8
|
|
if key_length > wire["limits"]["string_or_bytes_length_max"]:
|
|
raise ProtocolError("LIMIT_EXCEEDED", key_path)
|
|
if offset + key_length > end:
|
|
raise ProtocolError("LENGTH_MISMATCH", key_path)
|
|
raw_key = data[offset : offset + key_length]
|
|
offset += key_length
|
|
try:
|
|
key = raw_key.decode("utf-8")
|
|
except UnicodeDecodeError as exc:
|
|
raise ProtocolError("INVALID_UTF8", key_path) from exc
|
|
if "\x00" in key:
|
|
raise ProtocolError("NUL_STRING", key_path)
|
|
if unicodedata.normalize("NFC", key) != key:
|
|
raise ProtocolError("NON_NFC", key_path)
|
|
if previous is not None and raw_key <= previous:
|
|
code = "DUPLICATE_MAP_KEY" if raw_key == previous else "MAP_KEY_ORDER"
|
|
raise ProtocolError(code, key_path)
|
|
previous = raw_key
|
|
child_path = f"{path}[{index}]"
|
|
_need(data, offset, 8, end, child_path)
|
|
value_length = struct.unpack_from(">Q", data, offset)[0]
|
|
offset += 8
|
|
child_end = offset + value_length
|
|
if child_end > end:
|
|
raise ProtocolError("LENGTH_MISMATCH", child_path)
|
|
child, consumed = _decode_ast(data, offset, child_end, wire, depth + 1, child_path, True)
|
|
if consumed != child_end:
|
|
raise ProtocolError("LENGTH_MISMATCH", child_path)
|
|
result.append([key, child])
|
|
offset = child_end
|
|
if offset != end:
|
|
raise ProtocolError("LENGTH_MISMATCH", path)
|
|
return result
|
|
|
|
|
|
def _decode_record(
|
|
data: bytes, start: int, end: int, wire: dict[str, Any], depth: int, path: str
|
|
) -> list[Any]:
|
|
count, offset = _container_count(data, start, end, wire, path)
|
|
result = []
|
|
previous = 0
|
|
for index in range(count):
|
|
_need(data, offset, 12, end, path)
|
|
tag = struct.unpack_from(">I", data, offset)[0]
|
|
length = struct.unpack_from(">Q", data, offset + 4)[0]
|
|
offset += 12
|
|
if tag <= previous:
|
|
raise ProtocolError("FIELD_ORDER", path + f".{tag}")
|
|
previous = tag
|
|
child_end = offset + length
|
|
if child_end > end:
|
|
raise ProtocolError("LENGTH_MISMATCH", path + f".{tag}")
|
|
child, consumed = _decode_ast(data, offset, child_end, wire, depth + 1, path + f".{tag}", True)
|
|
if consumed != child_end:
|
|
raise ProtocolError("LENGTH_MISMATCH", path + f".{tag}")
|
|
result.append([tag, child])
|
|
offset = child_end
|
|
if offset != end:
|
|
raise ProtocolError("LENGTH_MISMATCH", path)
|
|
return result
|
|
|
|
|
|
def _read_value_file(path: pathlib.Path) -> Any:
|
|
return load_json(path)
|
|
|
|
|
|
def main(argv: list[str] | None = None) -> int:
|
|
default_schema = pathlib.Path(__file__).resolve().parent / "schema"
|
|
parser = argparse.ArgumentParser(description=__doc__)
|
|
sub = parser.add_subparsers(dest="command", required=True)
|
|
|
|
validate = sub.add_parser("validate")
|
|
validate.add_argument("--schema-dir", type=pathlib.Path, default=default_schema)
|
|
|
|
manifest = sub.add_parser("write-manifest")
|
|
manifest.add_argument("--schema-dir", type=pathlib.Path, default=default_schema)
|
|
|
|
emit_tables = sub.add_parser("emit-tables")
|
|
emit_tables.add_argument("--schema-dir", type=pathlib.Path, default=default_schema)
|
|
emit_tables.add_argument("--output", type=pathlib.Path, required=True)
|
|
|
|
check = sub.add_parser("check-output")
|
|
check.add_argument("--schema-dir", type=pathlib.Path, default=default_schema)
|
|
check.add_argument("--output", type=pathlib.Path, required=True)
|
|
|
|
schema_digest = sub.add_parser("schema-digest")
|
|
schema_digest.add_argument("--schema-dir", type=pathlib.Path, default=default_schema)
|
|
|
|
encode_value_parser = sub.add_parser("encode-value")
|
|
encode_value_parser.add_argument("--schema-dir", type=pathlib.Path, default=default_schema)
|
|
encode_value_parser.add_argument("--json", type=pathlib.Path, required=True)
|
|
|
|
encode_record_parser = sub.add_parser("encode-record")
|
|
encode_record_parser.add_argument("record")
|
|
encode_record_parser.add_argument("--schema-dir", type=pathlib.Path, default=default_schema)
|
|
encode_record_parser.add_argument("--json", type=pathlib.Path, required=True)
|
|
|
|
digest_parser = sub.add_parser("digest")
|
|
digest_parser.add_argument("domain")
|
|
digest_parser.add_argument("--schema-dir", type=pathlib.Path, default=default_schema)
|
|
digest_parser.add_argument("--json", type=pathlib.Path)
|
|
digest_parser.add_argument("--hex")
|
|
digest_parser.add_argument("--record")
|
|
digest_parser.add_argument("--kind", type=int)
|
|
digest_parser.add_argument("--record-schema", type=int, default=1)
|
|
|
|
args = parser.parse_args(argv)
|
|
try:
|
|
if args.command == "write-manifest":
|
|
SchemaBundle.from_dir(args.schema_dir, check_manifest=False)
|
|
write_schema_manifest(args.schema_dir)
|
|
return 0
|
|
if args.command == "emit-tables":
|
|
generate_table(args.schema_dir, args.output)
|
|
return 0
|
|
if args.command == "check-output":
|
|
expected = expected_table(args.schema_dir)
|
|
if not args.output.is_file() or args.output.read_bytes() != expected:
|
|
raise SchemaError(f"generated output is stale: {args.output}")
|
|
return 0
|
|
bundle = SchemaBundle.from_dir(args.schema_dir)
|
|
if args.command == "validate":
|
|
return 0
|
|
if args.command == "schema-digest":
|
|
print(bundle.schema_digest)
|
|
return 0
|
|
if args.command == "encode-value":
|
|
print(encode_envelope(_read_value_file(args.json), bundle.wire).hex())
|
|
return 0
|
|
if args.command == "encode-record":
|
|
print(bundle.encode_record(args.record, _read_value_file(args.json)).hex())
|
|
return 0
|
|
if args.command == "digest":
|
|
if args.hex is not None:
|
|
value: Any = bytes.fromhex(args.hex)
|
|
elif args.json is not None:
|
|
value = _read_value_file(args.json)
|
|
else:
|
|
raise SchemaError("digest requires --hex or --json")
|
|
print(
|
|
bundle.digest(
|
|
args.domain,
|
|
value,
|
|
record_name=args.record,
|
|
kind=args.kind,
|
|
schema=args.record_schema,
|
|
).hex()
|
|
)
|
|
return 0
|
|
except (SchemaError, ProtocolError, ValueError) as exc:
|
|
print(exc, file=sys.stderr)
|
|
return 2
|
|
return 2
|
|
|
|
|
|
if __name__ == "__main__":
|
|
raise SystemExit(main())
|