.beb and .json serialization of the document
This commit is contained in:
57
src/bedit_core/serialization/__init__.py
Normal file
57
src/bedit_core/serialization/__init__.py
Normal file
@@ -0,0 +1,57 @@
|
|||||||
|
"""Public file loading and saving API for bedit documents."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from pathlib import Path
|
||||||
|
|
||||||
|
from bedit_core.models import Document
|
||||||
|
|
||||||
|
from . import beb_codec, json_codec
|
||||||
|
from .migrations import migrate_document_layout, migrate_file_layout
|
||||||
|
from .schema import SerializationError, document_from_data, document_to_data
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"BEB_FILE_FORMAT_VERSION",
|
||||||
|
"JSON_FILE_FORMAT_VERSION",
|
||||||
|
"SerializationError",
|
||||||
|
"load",
|
||||||
|
"save",
|
||||||
|
]
|
||||||
|
|
||||||
|
BEB_FILE_FORMAT_VERSION = beb_codec.FILE_FORMAT_VERSION
|
||||||
|
JSON_FILE_FORMAT_VERSION = json_codec.FILE_FORMAT_VERSION
|
||||||
|
|
||||||
|
|
||||||
|
def load(path: str | Path) -> Document:
|
||||||
|
"""Load a document from ``path``.
|
||||||
|
|
||||||
|
The filename extension selects the JSON or BEB codec. File-layout
|
||||||
|
migrations and then document-layout migrations are applied before the
|
||||||
|
raw schema is converted to a :class:`~bedit_core.models.Document`.
|
||||||
|
"""
|
||||||
|
file_path = Path(path)
|
||||||
|
codec = _codec(file_path)
|
||||||
|
file_version, file_data = codec.load_data(file_path)
|
||||||
|
document_data = migrate_file_layout(
|
||||||
|
file_data,
|
||||||
|
codec.FORMAT_NAME,
|
||||||
|
file_version,
|
||||||
|
)
|
||||||
|
return document_from_data(migrate_document_layout(document_data))
|
||||||
|
|
||||||
|
|
||||||
|
def save(document: Document, path: str | Path) -> None:
|
||||||
|
"""Save ``document`` to a JSON or BEB file selected by ``path``'s suffix."""
|
||||||
|
file_path = Path(path)
|
||||||
|
codec = _codec(file_path)
|
||||||
|
codec.save_data(document_to_data(document), file_path)
|
||||||
|
|
||||||
|
|
||||||
|
def _codec(path: Path) -> Any:
|
||||||
|
"""Return the codec module associated with a supported filename suffix."""
|
||||||
|
suffix = path.suffix.lower()
|
||||||
|
if suffix == ".json":
|
||||||
|
return json_codec
|
||||||
|
if suffix == ".beb":
|
||||||
|
return beb_codec
|
||||||
|
raise ValueError(f"unsupported file extension {path.suffix!r}; expected '.json' or '.beb'")
|
||||||
75
src/bedit_core/serialization/beb_codec.py
Normal file
75
src/bedit_core/serialization/beb_codec.py
Normal file
@@ -0,0 +1,75 @@
|
|||||||
|
"""Versioned binary BEB codec.
|
||||||
|
|
||||||
|
The uncompressed header contains the magic bytes and codec version, allowing
|
||||||
|
the correct decoder to be selected even if later versions change compression
|
||||||
|
or serialization methods.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import zlib
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Mapping
|
||||||
|
|
||||||
|
import msgpack
|
||||||
|
|
||||||
|
FORMAT_NAME = "beb"
|
||||||
|
FILE_FORMAT_VERSION = 1
|
||||||
|
BEB_MAGIC = b"BEB\x00"
|
||||||
|
_VERSION_SIZE = 4
|
||||||
|
|
||||||
|
|
||||||
|
def load_data(path: Path) -> tuple[int, Mapping[str, Any]]:
|
||||||
|
"""Read a BEB header and decode its raw document mapping.
|
||||||
|
|
||||||
|
An unknown version is returned with an empty mapping so the migration
|
||||||
|
layer can produce the standard unsupported-version error without trying
|
||||||
|
an incompatible decoder.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
payload = path.read_bytes()
|
||||||
|
if not payload.startswith(BEB_MAGIC):
|
||||||
|
raise ValueError("missing BEB file header")
|
||||||
|
header_end = len(BEB_MAGIC) + _VERSION_SIZE
|
||||||
|
if len(payload) < header_end:
|
||||||
|
raise ValueError("truncated BEB file header")
|
||||||
|
version = int.from_bytes(payload[len(BEB_MAGIC) : header_end], "big")
|
||||||
|
encoded = payload[header_end:]
|
||||||
|
decoder = _DECODERS.get(version)
|
||||||
|
if decoder is None:
|
||||||
|
# The version can be inspected without trying to decompress or
|
||||||
|
# deserialize using the wrong algorithm.
|
||||||
|
return version, {}
|
||||||
|
data = decoder(encoded)
|
||||||
|
except (OSError, ValueError, zlib.error, msgpack.exceptions.MsgpackException) as exc:
|
||||||
|
raise ValueError(f"could not read BEB document {path}: {exc}") from exc
|
||||||
|
if not isinstance(data, Mapping):
|
||||||
|
raise ValueError(f"BEB document {path} must contain a map at its root")
|
||||||
|
return version, data
|
||||||
|
|
||||||
|
|
||||||
|
def save_data(data: Mapping[str, Any], path: Path) -> None:
|
||||||
|
"""Encode raw document data using the current BEB version and write it."""
|
||||||
|
encoded = _encode_v1(data)
|
||||||
|
version = FILE_FORMAT_VERSION.to_bytes(_VERSION_SIZE, "big")
|
||||||
|
path.write_bytes(BEB_MAGIC + version + encoded)
|
||||||
|
|
||||||
|
|
||||||
|
def _encode_v1(data: Mapping[str, Any]) -> bytes:
|
||||||
|
"""Encode BEB version 1 as zlib-compressed MessagePack."""
|
||||||
|
return zlib.compress(msgpack.packb(dict(data), use_bin_type=True))
|
||||||
|
|
||||||
|
|
||||||
|
def _decode_v1(payload: bytes) -> Any:
|
||||||
|
"""Decode a zlib-compressed MessagePack BEB version 1 payload."""
|
||||||
|
return msgpack.unpackb(
|
||||||
|
zlib.decompress(payload),
|
||||||
|
raw=False,
|
||||||
|
strict_map_key=False,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
# Keep old decoders when adding a new BEB encoding version.
|
||||||
|
_DECODERS = {
|
||||||
|
1: _decode_v1,
|
||||||
|
}
|
||||||
36
src/bedit_core/serialization/json_codec.py
Normal file
36
src/bedit_core/serialization/json_codec.py
Normal file
@@ -0,0 +1,36 @@
|
|||||||
|
"""Versioned UTF-8 JSON codec used by the serialization API."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import json
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any, Mapping
|
||||||
|
|
||||||
|
FORMAT_NAME = "json"
|
||||||
|
FILE_FORMAT_VERSION = 1
|
||||||
|
|
||||||
|
|
||||||
|
def load_data(path: Path) -> tuple[int, Mapping[str, Any]]:
|
||||||
|
"""Read JSON and return its file-layout version and raw root object.
|
||||||
|
|
||||||
|
JSON files without ``file_format_version`` are interpreted as version 1.
|
||||||
|
Schema conversion and migrations are intentionally handled by higher
|
||||||
|
layers.
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
with path.open("r", encoding="utf-8") as stream:
|
||||||
|
data = json.load(stream)
|
||||||
|
except (OSError, UnicodeError, json.JSONDecodeError) as exc:
|
||||||
|
raise ValueError(f"could not read JSON document {path}: {exc}") from exc
|
||||||
|
if not isinstance(data, Mapping):
|
||||||
|
raise ValueError(f"JSON document {path} must contain an object at its root")
|
||||||
|
version = data.get("file_format_version", 1)
|
||||||
|
return version, data
|
||||||
|
|
||||||
|
|
||||||
|
def save_data(data: Mapping[str, Any], path: Path) -> None:
|
||||||
|
"""Write raw document data as indented UTF-8 JSON at the current version."""
|
||||||
|
output = {"file_format_version": FILE_FORMAT_VERSION, **data}
|
||||||
|
with path.open("w", encoding="utf-8") as stream:
|
||||||
|
json.dump(output, stream, ensure_ascii=False, indent=2)
|
||||||
|
stream.write("\n")
|
||||||
104
src/bedit_core/serialization/migrations.py
Normal file
104
src/bedit_core/serialization/migrations.py
Normal file
@@ -0,0 +1,104 @@
|
|||||||
|
"""Registries and runners for file-layout and document-layout migrations.
|
||||||
|
|
||||||
|
Register a migration under version ``N`` to convert raw data from ``N`` to
|
||||||
|
``N + 1``. JSON and BEB file migrations are independent; document migrations
|
||||||
|
run after file migrations and are shared by both codecs.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Callable, Mapping
|
||||||
|
from copy import deepcopy
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from . import beb_codec, json_codec
|
||||||
|
from .schema import SerializationError
|
||||||
|
|
||||||
|
DOCUMENT_FORMAT_VERSION = 1
|
||||||
|
|
||||||
|
Migration = Callable[[dict[str, Any]], dict[str, Any]]
|
||||||
|
|
||||||
|
# A migration registered under N converts version N to N + 1. File migrations
|
||||||
|
# are separate because JSON and BEB may evolve independently.
|
||||||
|
FILE_FORMAT_VERSIONS = {
|
||||||
|
json_codec.FORMAT_NAME: json_codec.FILE_FORMAT_VERSION,
|
||||||
|
beb_codec.FORMAT_NAME: beb_codec.FILE_FORMAT_VERSION,
|
||||||
|
}
|
||||||
|
FILE_MIGRATIONS: dict[str, dict[int, Migration]] = {
|
||||||
|
json_codec.FORMAT_NAME: {},
|
||||||
|
beb_codec.FORMAT_NAME: {},
|
||||||
|
}
|
||||||
|
DOCUMENT_MIGRATIONS: dict[int, Migration] = {}
|
||||||
|
|
||||||
|
|
||||||
|
def migrate_file_layout(
|
||||||
|
data: Mapping[str, Any],
|
||||||
|
file_format: str,
|
||||||
|
version: int,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Migrate codec-specific raw data to that codec's current file layout.
|
||||||
|
|
||||||
|
The input is copied, so migrations cannot mutate the codec's decoded
|
||||||
|
object. ``file_format_version`` is removed before document migration.
|
||||||
|
"""
|
||||||
|
migrated = deepcopy(dict(data))
|
||||||
|
version = _version(version, f"{file_format} file format version")
|
||||||
|
try:
|
||||||
|
target = FILE_FORMAT_VERSIONS[file_format]
|
||||||
|
migrations = FILE_MIGRATIONS[file_format]
|
||||||
|
except KeyError:
|
||||||
|
raise SerializationError(f"unknown file format {file_format!r}") from None
|
||||||
|
migrated = _apply(
|
||||||
|
migrated,
|
||||||
|
version,
|
||||||
|
target,
|
||||||
|
migrations,
|
||||||
|
f"{file_format} file",
|
||||||
|
)
|
||||||
|
migrated.pop("file_format_version", None)
|
||||||
|
return migrated
|
||||||
|
|
||||||
|
|
||||||
|
def migrate_document_layout(data: Mapping[str, Any]) -> dict[str, Any]:
|
||||||
|
"""Migrate raw document fields to ``DOCUMENT_FORMAT_VERSION``."""
|
||||||
|
migrated = deepcopy(dict(data))
|
||||||
|
version = _version(migrated.get("format_version"), "format_version")
|
||||||
|
migrated = _apply(
|
||||||
|
migrated,
|
||||||
|
version,
|
||||||
|
DOCUMENT_FORMAT_VERSION,
|
||||||
|
DOCUMENT_MIGRATIONS,
|
||||||
|
"document",
|
||||||
|
)
|
||||||
|
migrated["format_version"] = DOCUMENT_FORMAT_VERSION
|
||||||
|
return migrated
|
||||||
|
|
||||||
|
|
||||||
|
def _apply(
|
||||||
|
data: dict[str, Any],
|
||||||
|
version: int,
|
||||||
|
target: int,
|
||||||
|
migrations: Mapping[int, Migration],
|
||||||
|
kind: str,
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""Apply consecutive migrations from ``version`` up to ``target``."""
|
||||||
|
if version > target:
|
||||||
|
raise SerializationError(
|
||||||
|
f"unsupported {kind} format version {version}; newest supported version is {target}"
|
||||||
|
)
|
||||||
|
while version < target:
|
||||||
|
migration = migrations.get(version)
|
||||||
|
if migration is None:
|
||||||
|
raise SerializationError(
|
||||||
|
f"no {kind} migration is registered from version {version} to {version + 1}"
|
||||||
|
)
|
||||||
|
data = migration(data)
|
||||||
|
version += 1
|
||||||
|
return data
|
||||||
|
|
||||||
|
|
||||||
|
def _version(value: Any, field: str) -> int:
|
||||||
|
"""Validate and return a positive integer version value."""
|
||||||
|
if not isinstance(value, int) or isinstance(value, bool) or value < 1:
|
||||||
|
raise SerializationError(f"{field}: expected a positive integer")
|
||||||
|
return value
|
||||||
381
src/bedit_core/serialization/schema.py
Normal file
381
src/bedit_core/serialization/schema.py
Normal file
@@ -0,0 +1,381 @@
|
|||||||
|
"""Conversion between raw serialization values and model dataclasses."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
from collections.abc import Mapping
|
||||||
|
from enum import Enum
|
||||||
|
from typing import Any
|
||||||
|
|
||||||
|
from bedit_core.models import (
|
||||||
|
BondCausality,
|
||||||
|
BondConnection,
|
||||||
|
BondPort,
|
||||||
|
Component,
|
||||||
|
ComponentID,
|
||||||
|
Connection,
|
||||||
|
ConnectionID,
|
||||||
|
Document,
|
||||||
|
EquationImplementation,
|
||||||
|
Graph,
|
||||||
|
GraphImplementation,
|
||||||
|
ID,
|
||||||
|
Interface,
|
||||||
|
Parameter,
|
||||||
|
ParameterID,
|
||||||
|
Port,
|
||||||
|
PortCausality,
|
||||||
|
PortID,
|
||||||
|
SignalConnection,
|
||||||
|
SignalDirection,
|
||||||
|
SignalPort,
|
||||||
|
)
|
||||||
|
|
||||||
|
class SerializationError(ValueError):
|
||||||
|
"""Raised when serialized data does not match the bedit schema."""
|
||||||
|
|
||||||
|
|
||||||
|
def document_to_data(document: Document) -> dict[str, Any]:
|
||||||
|
"""Convert a document tree to JSON/MessagePack-compatible values."""
|
||||||
|
if not isinstance(document, Document):
|
||||||
|
raise TypeError("document must be a Document")
|
||||||
|
return {
|
||||||
|
"format_version": document.format_version,
|
||||||
|
"id": str(document.id),
|
||||||
|
"name": document.name,
|
||||||
|
"root": {
|
||||||
|
str(component_id): _component_to_data(component)
|
||||||
|
for component_id, component in document.root.items()
|
||||||
|
},
|
||||||
|
"metadata": _plain_value(document.metadata),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def document_from_data(data: Mapping[str, Any]) -> Document:
|
||||||
|
"""Validate raw values and construct a complete document dataclass tree."""
|
||||||
|
obj = _mapping(data, "document")
|
||||||
|
return Document(
|
||||||
|
format_version=_integer(_required(obj, "format_version", "document"), "document.format_version"),
|
||||||
|
id=ID(_string(_required(obj, "id", "document"), "document.id")),
|
||||||
|
name=_string(_required(obj, "name", "document"), "document.name"),
|
||||||
|
root={
|
||||||
|
ComponentID(_string(key, "document.root key")): _component_from_data(value, f"document.root[{key!r}]")
|
||||||
|
for key, value in _mapping(_required(obj, "root", "document"), "document.root").items()
|
||||||
|
},
|
||||||
|
metadata=_metadata(obj.get("metadata")),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _component_to_data(component: Component) -> dict[str, Any]:
|
||||||
|
"""Convert a component and its nested model objects to raw values."""
|
||||||
|
return {
|
||||||
|
"name": component.name,
|
||||||
|
"interface": {
|
||||||
|
"ports": {
|
||||||
|
str(port_id): _port_to_data(port)
|
||||||
|
for port_id, port in component.interface.ports.items()
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"parameters": {
|
||||||
|
str(parameter_id): _parameter_to_data(parameter)
|
||||||
|
for parameter_id, parameter in component.parameters.items()
|
||||||
|
},
|
||||||
|
"implementation": _implementation_to_data(component.implementation),
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _component_from_data(value: Any, where: str) -> Component:
|
||||||
|
"""Construct a component from raw values, using ``where`` in errors."""
|
||||||
|
obj = _mapping(value, where)
|
||||||
|
interface = _mapping(_required(obj, "interface", where), f"{where}.interface")
|
||||||
|
ports = _mapping(_required(interface, "ports", f"{where}.interface"), f"{where}.interface.ports")
|
||||||
|
parameters = _mapping(_required(obj, "parameters", where), f"{where}.parameters")
|
||||||
|
return Component(
|
||||||
|
name=_string(_required(obj, "name", where), f"{where}.name"),
|
||||||
|
interface=Interface(
|
||||||
|
ports={
|
||||||
|
PortID(_string(key, f"{where}.interface.ports key")): _port_from_data(port, f"{where}.interface.ports[{key!r}]")
|
||||||
|
for key, port in ports.items()
|
||||||
|
}
|
||||||
|
),
|
||||||
|
parameters={
|
||||||
|
ParameterID(_string(key, f"{where}.parameters key")): _parameter_from_data(parameter, f"{where}.parameters[{key!r}]")
|
||||||
|
for key, parameter in parameters.items()
|
||||||
|
},
|
||||||
|
implementation=_implementation_from_data(_required(obj, "implementation", where), f"{where}.implementation"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _port_to_data(port: Port) -> dict[str, Any]:
|
||||||
|
"""Convert a derived port and add its serialized type discriminator."""
|
||||||
|
common = {
|
||||||
|
"name": port.name,
|
||||||
|
"direction": port.direction.value,
|
||||||
|
"multiplicity": port.multiplicity,
|
||||||
|
"matrix_size": list(port.matrix_size),
|
||||||
|
"description": port.description,
|
||||||
|
}
|
||||||
|
if isinstance(port, SignalPort):
|
||||||
|
return {
|
||||||
|
"port_type": "signal",
|
||||||
|
**common,
|
||||||
|
"value_type": port.value_type,
|
||||||
|
"quantity": port.quantity,
|
||||||
|
"unit": port.unit,
|
||||||
|
}
|
||||||
|
if isinstance(port, BondPort):
|
||||||
|
return {
|
||||||
|
"port_type": "bond",
|
||||||
|
**common,
|
||||||
|
"domain": port.domain,
|
||||||
|
"causality_preference": port.causality_preference.value,
|
||||||
|
}
|
||||||
|
raise TypeError(f"unsupported port class: {type(port).__name__}")
|
||||||
|
|
||||||
|
|
||||||
|
def _port_from_data(value: Any, where: str) -> Port:
|
||||||
|
"""Construct the port subclass selected by ``port_type``."""
|
||||||
|
obj = _mapping(value, where)
|
||||||
|
port_type = _string(_required(obj, "port_type", where), f"{where}.port_type")
|
||||||
|
common = {
|
||||||
|
"name": _string(_required(obj, "name", where), f"{where}.name"),
|
||||||
|
"direction": _enum(SignalDirection, _required(obj, "direction", where), f"{where}.direction"),
|
||||||
|
"multiplicity": _boolean(obj.get("multiplicity", False), f"{where}.multiplicity"),
|
||||||
|
"matrix_size": _matrix_size(obj.get("matrix_size", [1, 1]), f"{where}.matrix_size"),
|
||||||
|
"description": _optional_string(obj.get("description"), f"{where}.description"),
|
||||||
|
}
|
||||||
|
if port_type == "signal":
|
||||||
|
return SignalPort(
|
||||||
|
**common,
|
||||||
|
value_type=_string(obj.get("value_type", "Real"), f"{where}.value_type"),
|
||||||
|
quantity=_optional_string(obj.get("quantity"), f"{where}.quantity"),
|
||||||
|
unit=_optional_string(obj.get("unit"), f"{where}.unit"),
|
||||||
|
)
|
||||||
|
if port_type == "bond":
|
||||||
|
return BondPort(
|
||||||
|
**common,
|
||||||
|
domain=_string(obj.get("domain", ""), f"{where}.domain"),
|
||||||
|
causality_preference=_enum(
|
||||||
|
PortCausality,
|
||||||
|
obj.get("causality_preference", PortCausality.INDIFFERENT.value),
|
||||||
|
f"{where}.causality_preference",
|
||||||
|
),
|
||||||
|
)
|
||||||
|
raise SerializationError(f"{where}.port_type: unsupported value {port_type!r}")
|
||||||
|
|
||||||
|
|
||||||
|
def _parameter_to_data(parameter: Parameter) -> dict[str, Any]:
|
||||||
|
"""Convert a parameter to raw values."""
|
||||||
|
return {
|
||||||
|
"name": parameter.name,
|
||||||
|
"value": _plain_value(parameter.value),
|
||||||
|
"value_type": parameter.value_type,
|
||||||
|
"quantity": parameter.quantity,
|
||||||
|
"unit": parameter.unit,
|
||||||
|
"description": parameter.description,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _parameter_from_data(value: Any, where: str) -> Parameter:
|
||||||
|
"""Construct a parameter from validated raw values."""
|
||||||
|
obj = _mapping(value, where)
|
||||||
|
return Parameter(
|
||||||
|
name=_string(_required(obj, "name", where), f"{where}.name"),
|
||||||
|
value=obj.get("value", 1.0),
|
||||||
|
value_type=_string(obj.get("value_type", "Real"), f"{where}.value_type"),
|
||||||
|
quantity=_optional_string(obj.get("quantity"), f"{where}.quantity"),
|
||||||
|
unit=_optional_string(obj.get("unit"), f"{where}.unit"),
|
||||||
|
description=_optional_string(obj.get("description"), f"{where}.description"),
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _implementation_to_data(implementation: GraphImplementation | EquationImplementation) -> dict[str, Any]:
|
||||||
|
"""Convert an implementation and add its serialized discriminator."""
|
||||||
|
if isinstance(implementation, GraphImplementation):
|
||||||
|
return {
|
||||||
|
"implementation_type": "graph",
|
||||||
|
"graph": _graph_to_data(implementation.graph),
|
||||||
|
}
|
||||||
|
if isinstance(implementation, EquationImplementation):
|
||||||
|
return {
|
||||||
|
"implementation_type": "equation",
|
||||||
|
"declarations": list(implementation.declarations),
|
||||||
|
"initial_equations": list(implementation.initial_equations),
|
||||||
|
"equations": list(implementation.equations),
|
||||||
|
}
|
||||||
|
raise TypeError(f"unsupported implementation class: {type(implementation).__name__}")
|
||||||
|
|
||||||
|
|
||||||
|
def _implementation_from_data(value: Any, where: str) -> GraphImplementation | EquationImplementation:
|
||||||
|
"""Construct the implementation selected by ``implementation_type``."""
|
||||||
|
obj = _mapping(value, where)
|
||||||
|
implementation_type = _string(
|
||||||
|
_required(obj, "implementation_type", where),
|
||||||
|
f"{where}.implementation_type",
|
||||||
|
)
|
||||||
|
if implementation_type == "graph":
|
||||||
|
return GraphImplementation(_graph_from_data(_required(obj, "graph", where), f"{where}.graph"))
|
||||||
|
if implementation_type == "equation":
|
||||||
|
return EquationImplementation(
|
||||||
|
declarations=_string_list(obj.get("declarations", []), f"{where}.declarations"),
|
||||||
|
initial_equations=_string_list(obj.get("initial_equations", []), f"{where}.initial_equations"),
|
||||||
|
equations=_string_list(obj.get("equations", []), f"{where}.equations"),
|
||||||
|
)
|
||||||
|
raise SerializationError(
|
||||||
|
f"{where}.implementation_type: unsupported value {implementation_type!r}"
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _graph_to_data(graph: Graph) -> dict[str, Any]:
|
||||||
|
"""Convert a graph, including its components and connections."""
|
||||||
|
return {
|
||||||
|
"components": {
|
||||||
|
str(component_id): _component_to_data(component)
|
||||||
|
for component_id, component in graph.components.items()
|
||||||
|
},
|
||||||
|
"connections": {
|
||||||
|
str(connection_id): _connection_to_data(connection)
|
||||||
|
for connection_id, connection in graph.connections.items()
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _graph_from_data(value: Any, where: str) -> Graph:
|
||||||
|
"""Construct a graph and all nested dataclasses from raw values."""
|
||||||
|
obj = _mapping(value, where)
|
||||||
|
components = _mapping(obj.get("components", {}), f"{where}.components")
|
||||||
|
connections = _mapping(obj.get("connections", {}), f"{where}.connections")
|
||||||
|
return Graph(
|
||||||
|
components={
|
||||||
|
ComponentID(_string(key, f"{where}.components key")): _component_from_data(component, f"{where}.components[{key!r}]")
|
||||||
|
for key, component in components.items()
|
||||||
|
},
|
||||||
|
connections={
|
||||||
|
ConnectionID(_string(key, f"{where}.connections key")): _connection_from_data(connection, f"{where}.connections[{key!r}]")
|
||||||
|
for key, connection in connections.items()
|
||||||
|
},
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def _connection_to_data(connection: Connection) -> dict[str, Any]:
|
||||||
|
"""Convert a derived connection and add its serialized discriminator."""
|
||||||
|
common = {"source": str(connection.source), "target": str(connection.target)}
|
||||||
|
if isinstance(connection, BondConnection):
|
||||||
|
return {
|
||||||
|
"connection_type": "bond",
|
||||||
|
**common,
|
||||||
|
"causality": connection.causality.value,
|
||||||
|
"undesired": connection.undesired,
|
||||||
|
}
|
||||||
|
if isinstance(connection, SignalConnection):
|
||||||
|
return {"connection_type": "signal", **common}
|
||||||
|
raise TypeError(f"unsupported connection class: {type(connection).__name__}")
|
||||||
|
|
||||||
|
|
||||||
|
def _connection_from_data(value: Any, where: str) -> Connection:
|
||||||
|
"""Construct the connection subclass selected by ``connection_type``."""
|
||||||
|
obj = _mapping(value, where)
|
||||||
|
# Files produced before the serializer existed had no connection discriminator.
|
||||||
|
connection_type = obj.get("connection_type")
|
||||||
|
if connection_type is None:
|
||||||
|
connection_type = "bond" if "causality" in obj or "undesired" in obj else "signal"
|
||||||
|
connection_type = _string(connection_type, f"{where}.connection_type")
|
||||||
|
common = {
|
||||||
|
"source": PortID(_string(_required(obj, "source", where), f"{where}.source")),
|
||||||
|
"target": PortID(_string(_required(obj, "target", where), f"{where}.target")),
|
||||||
|
}
|
||||||
|
if connection_type == "signal":
|
||||||
|
return SignalConnection(**common)
|
||||||
|
if connection_type == "bond":
|
||||||
|
return BondConnection(
|
||||||
|
**common,
|
||||||
|
causality=_enum(BondCausality, obj.get("causality", "none"), f"{where}.causality"),
|
||||||
|
undesired=_boolean(obj.get("undesired", False), f"{where}.undesired"),
|
||||||
|
)
|
||||||
|
raise SerializationError(f"{where}.connection_type: unsupported value {connection_type!r}")
|
||||||
|
|
||||||
|
|
||||||
|
def _plain_value(value: Any) -> Any:
|
||||||
|
"""Recursively convert metadata or parameter values to codec-safe values."""
|
||||||
|
if isinstance(value, Enum):
|
||||||
|
return value.value
|
||||||
|
if value is None or isinstance(value, (str, int, float, bool)):
|
||||||
|
return value
|
||||||
|
if isinstance(value, Mapping):
|
||||||
|
return {str(key): _plain_value(item) for key, item in value.items()}
|
||||||
|
if isinstance(value, (list, tuple)):
|
||||||
|
return [_plain_value(item) for item in value]
|
||||||
|
raise TypeError(f"value of type {type(value).__name__} is not serializable")
|
||||||
|
|
||||||
|
|
||||||
|
def _required(obj: Mapping[str, Any], key: str, where: str) -> Any:
|
||||||
|
"""Return a required mapping value or raise a contextual schema error."""
|
||||||
|
if key not in obj:
|
||||||
|
raise SerializationError(f"{where}: missing required field {key!r}")
|
||||||
|
return obj[key]
|
||||||
|
|
||||||
|
|
||||||
|
def _mapping(value: Any, where: str) -> Mapping[str, Any]:
|
||||||
|
"""Validate that ``value`` is a mapping."""
|
||||||
|
if not isinstance(value, Mapping):
|
||||||
|
raise SerializationError(f"{where}: expected an object")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _string(value: Any, where: str) -> str:
|
||||||
|
"""Validate that ``value`` is a string."""
|
||||||
|
if not isinstance(value, str):
|
||||||
|
raise SerializationError(f"{where}: expected a string")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _integer(value: Any, where: str) -> int:
|
||||||
|
"""Validate that ``value`` is an integer but not a boolean."""
|
||||||
|
if not isinstance(value, int) or isinstance(value, bool):
|
||||||
|
raise SerializationError(f"{where}: expected an integer")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _boolean(value: Any, where: str) -> bool:
|
||||||
|
"""Validate that ``value`` is a boolean."""
|
||||||
|
if not isinstance(value, bool):
|
||||||
|
raise SerializationError(f"{where}: expected a boolean")
|
||||||
|
return value
|
||||||
|
|
||||||
|
|
||||||
|
def _optional_string(value: Any, where: str) -> str | None:
|
||||||
|
"""Validate that ``value`` is either a string or ``None``."""
|
||||||
|
return None if value is None else _string(value, where)
|
||||||
|
|
||||||
|
|
||||||
|
def _string_list(value: Any, where: str) -> list[str]:
|
||||||
|
"""Validate and copy a list of strings."""
|
||||||
|
if not isinstance(value, list):
|
||||||
|
raise SerializationError(f"{where}: expected a list")
|
||||||
|
return [_string(item, f"{where}[{index}]") for index, item in enumerate(value)]
|
||||||
|
|
||||||
|
|
||||||
|
def _matrix_size(value: Any, where: str) -> list[int]:
|
||||||
|
"""Validate and copy a two-integer matrix size."""
|
||||||
|
if (
|
||||||
|
not isinstance(value, list)
|
||||||
|
or len(value) != 2
|
||||||
|
or any(not isinstance(item, int) or isinstance(item, bool) for item in value)
|
||||||
|
):
|
||||||
|
raise SerializationError(f"{where}: expected a list of two integers")
|
||||||
|
return list(value)
|
||||||
|
|
||||||
|
|
||||||
|
def _metadata(value: Any) -> dict[str, Any] | None:
|
||||||
|
"""Validate and copy optional document metadata."""
|
||||||
|
if value is None:
|
||||||
|
return None
|
||||||
|
return dict(_mapping(value, "document.metadata"))
|
||||||
|
|
||||||
|
|
||||||
|
def _enum(enum_type: type[Enum], value: Any, where: str) -> Any:
|
||||||
|
"""Convert a serialized enum value or raise a contextual schema error."""
|
||||||
|
try:
|
||||||
|
return enum_type(value)
|
||||||
|
except (TypeError, ValueError):
|
||||||
|
choices = ", ".join(repr(member.value) for member in enum_type)
|
||||||
|
raise SerializationError(f"{where}: expected one of {choices}") from None
|
||||||
37
tests/conftest.py
Normal file
37
tests/conftest.py
Normal file
@@ -0,0 +1,37 @@
|
|||||||
|
"""Shared pytest fixtures for bedit tests.
|
||||||
|
|
||||||
|
Qt-specific fixtures can be added here later when GUI testing starts. Keeping
|
||||||
|
the model fixtures independent from Qt lets the core suite stay lightweight.
|
||||||
|
"""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from bedit_core.models import (
|
||||||
|
Component,
|
||||||
|
ComponentID,
|
||||||
|
Document,
|
||||||
|
Graph,
|
||||||
|
GraphImplementation,
|
||||||
|
ID,
|
||||||
|
Interface,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def minimal_document() -> Document:
|
||||||
|
"""Return the smallest useful graph document for core tests."""
|
||||||
|
root_id = ComponentID("root")
|
||||||
|
root = Component(
|
||||||
|
name="Root",
|
||||||
|
interface=Interface(),
|
||||||
|
parameters={},
|
||||||
|
implementation=GraphImplementation(Graph()),
|
||||||
|
)
|
||||||
|
return Document(
|
||||||
|
format_version=1,
|
||||||
|
id=ID("document"),
|
||||||
|
name="Test document",
|
||||||
|
root={root_id: root},
|
||||||
|
)
|
||||||
14
tests/unit/test_models.py
Normal file
14
tests/unit/test_models.py
Normal file
@@ -0,0 +1,14 @@
|
|||||||
|
"""Basic smoke tests for the core model."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from bedit_core.models import Document, GraphImplementation
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
def test_minimal_document_has_graph_root(minimal_document: Document) -> None:
|
||||||
|
root = next(iter(minimal_document.root.values()))
|
||||||
|
|
||||||
|
assert isinstance(root.implementation, GraphImplementation)
|
||||||
17
tests/unit/test_serialization.py
Normal file
17
tests/unit/test_serialization.py
Normal file
@@ -0,0 +1,17 @@
|
|||||||
|
"""Basic smoke tests for document serialization."""
|
||||||
|
|
||||||
|
from __future__ import annotations
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from bedit_core.models import Document
|
||||||
|
from bedit_core.serialization import load, save
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.unit
|
||||||
|
def test_json_round_trip(tmp_path, minimal_document: Document) -> None:
|
||||||
|
path = tmp_path / "document.json"
|
||||||
|
|
||||||
|
save(minimal_document, path)
|
||||||
|
|
||||||
|
assert load(path) == minimal_document
|
||||||
Reference in New Issue
Block a user