From d0c285658ce787a35da897823b3443e725d41681 Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Thu, 17 Sep 2026 22:28:34 +0200 Subject: [PATCH 01/14] Add the config base and the model metadata record to mace-core (CORE-2, #1556) `mace_core.config.ReforgeBaseConfig` is the pydantic base every v1 config schema derives from, and `mace_core.metadata.ModelMetadata` the versioned record every trained model carries. Both are framework-free; a test in a fresh interpreter asserts neither imports torch, jax or e3nn. Config: - A schema is a tree of `ConfigSection`s; `load()` reads one TOML, YAML or JSON file and applies dotted CLI overrides (`--model.radial.cutoff 5.0`, `--seed=7`), precedence defaults < file < CLI. Nothing else feeds a config: no environment variables, no dotenv files. - Overrides are parsed by a plain tokenizer, not pydantic-settings or argparse. A value starting with `[` or `{`, or `null`, is JSON, so a whole section, list or dict can be given. Each override merges onto the file values in order by one rule: dict-valued fields gain or replace entries, list-valued fields are replaced whole. - Unknown keys are `ConfigError`s that name the key by its dotted path and the nearest valid neighbour: CLI keys are checked against the schema's dotted paths before anything is built, file keys and keys inside a JSON-valued override by pydantic (`extra="forbid"` on every level). Wrong types re-raise pydantic's `ValidationError`. - `to_resolved_dict()` exports every field, defaults filled in, in schema order, as a fixed point: loading it back and resolving again gives the same dict. `to_user_dict()` exports only what the file and CLI set. - Field shapes the contract cannot keep are rejected at class definition: sets (order is not stable across runs), aliases and computed fields (the export would not validate back), a union of section classes (a value must not become whichever alternative accepts it; offer alternatives as one optional field each), and sections that do not subclass `ConfigSection`. Metadata: - `ModelMetadata` holds the config as written and as resolved (`ConfigRecord.from_config`), provenance (code version, git commit), one `DataSourceSummary` per data source with the `_` reference-key convention, `E0Details` per head (explicit or estimated, with method and parameters), the model's DOI, structured citations, notes, and `parents`: the models it was built from, each with a role (`initial_weights` for fine-tuning and continued training, `teacher` for distillation) and its own record embedded, so the lineage travels with the model. - `to_json()` refuses a record that would not read back equal, so a tuple, datetime or NaN is never stored as something else; infinity round-trips. `from_json()` checks `schema_version` before any field is read and rejects a newer version with a message naming both versions. - `format_citations()` renders a numbered block for printing. Floors: Python 3.10, pydantic 2.7; the suite passes at both floors and at current releases. Tests live in `test_mace_core_config.py` and `test_mace_core_metadata.py`. --- packages/mace-core/README.md | 17 +- packages/mace-core/pyproject.toml | 8 +- packages/mace-core/src/mace_core/__init__.py | 15 +- .../src/mace_core/config/__init__.py | 14 + .../mace-core/src/mace_core/config/base.py | 368 ++++++++++++++ packages/mace-core/src/mace_core/metadata.py | 237 +++++++++ .../mace-core/tests/test_mace_core_config.py | 475 ++++++++++++++++++ .../tests/test_mace_core_metadata.py | 234 +++++++++ 8 files changed, 1364 insertions(+), 4 deletions(-) create mode 100644 packages/mace-core/src/mace_core/config/__init__.py create mode 100644 packages/mace-core/src/mace_core/config/base.py create mode 100644 packages/mace-core/src/mace_core/metadata.py create mode 100644 packages/mace-core/tests/test_mace_core_config.py create mode 100644 packages/mace-core/tests/test_mace_core_metadata.py diff --git a/packages/mace-core/README.md b/packages/mace-core/README.md index 0b2cb598b..f8f6999b2 100644 --- a/packages/mace-core/README.md +++ b/packages/mace-core/README.md @@ -5,4 +5,19 @@ observables, the kernel-backend protocol, Clebsch-Gordan, neighbours and the data spec. It imports no torch, no jax and no e3nn, so both implementation packages can depend on it without depending on each other. -Distribution `mace-core`, import name `mace_core`. Scaffold only for now. +Distribution `mace-core`, import name `mace_core`. + +What is here so far: + +- `mace_core.config` — `ReforgeBaseConfig`, the pydantic base every v1 + config schema derives from: one TOML/YAML/JSON file plus dotted CLI + overrides (`--model.num_interactions 3`), + precedence defaults < file < CLI, + unknown keys rejected with the nearest valid neighbour named, and + `to_resolved_dict()` for the fully defaulted, round-trippable export. +- `mace_core.metadata` — `ModelMetadata`, the versioned record every trained + model carries (config as written and as resolved, provenance, a summary per + data source, E0 details per head, DOI, citations, notes, and the records of + the models it was fine-tuned or distilled from), with a JSON round trip, + `ConfigRecord.from_config()` to embed a config in its fixed-point form, and + `format_citations()`. diff --git a/packages/mace-core/pyproject.toml b/packages/mace-core/pyproject.toml index 7caf4c664..487c094a2 100644 --- a/packages/mace-core/pyproject.toml +++ b/packages/mace-core/pyproject.toml @@ -17,7 +17,13 @@ classifiers = [ "Programming Language :: Python :: 3.13", "Operating System :: OS Independent", ] -dependencies = [] +dependencies = [ + # The suite passes at this floor (2.7.4 on Python 3.10) and at current + # releases. + "pydantic>=2.7", + "pyyaml>=6.0", + "tomli>=2.0", +] [project.urls] Homepage = "https://github.com/ACEsuit/mace" diff --git a/packages/mace-core/src/mace_core/__init__.py b/packages/mace-core/src/mace_core/__init__.py index 4122cd7d8..882339eb6 100644 --- a/packages/mace-core/src/mace_core/__init__.py +++ b/packages/mace-core/src/mace_core/__init__.py @@ -1,11 +1,22 @@ """Framework-agnostic contract and pure math for MACE v1. -Scaffold only. The public surface arrives with the tickets that build it. +Re-exports the public surface of the lightweight submodules. Heavy ones (the +Clebsch-Gordan basis, neighbours) are imported from their own modules. """ from importlib.metadata import PackageNotFoundError, version -__all__ = ["__version__"] +from mace_core.config import ConfigError, ConfigSection, ReforgeBaseConfig +from mace_core.metadata import ModelMetadata, format_citations + +__all__ = [ + "ConfigError", + "ConfigSection", + "ModelMetadata", + "ReforgeBaseConfig", + "__version__", + "format_citations", +] #: Version of the installed `mace-core` distribution. Read from installed metadata #: rather than hardcoded, so it cannot drift from what pip resolved. diff --git a/packages/mace-core/src/mace_core/config/__init__.py b/packages/mace-core/src/mace_core/config/__init__.py new file mode 100644 index 000000000..9d6eaf542 --- /dev/null +++ b/packages/mace-core/src/mace_core/config/__init__.py @@ -0,0 +1,14 @@ +"""Configuration schemas for MACE v1. + +The base machinery lives in `base`; the training schema sections (model, data, +training, ...) arrive with their own tickets and are re-exported from here. +""" + +from mace_core.config.base import ( + ConfigError, + ConfigSection, + ReforgeBaseConfig, + read_config_file, +) + +__all__ = ["ConfigError", "ConfigSection", "ReforgeBaseConfig", "read_config_file"] diff --git a/packages/mace-core/src/mace_core/config/base.py b/packages/mace-core/src/mace_core/config/base.py new file mode 100644 index 000000000..e9ebf81c4 --- /dev/null +++ b/packages/mace-core/src/mace_core/config/base.py @@ -0,0 +1,368 @@ +"""The base class every v1 configuration schema derives from. + +A configuration is a tree of fields. A field holds either one value (`seed`, +`cutoff`) or a named group of further fields; such a group is a *section*. +The root of the tree subclasses `ReforgeBaseConfig`, every section +subclasses `ConfigSection`: + + class RadialSection(ConfigSection): + num_bessel: int = 8 + cutoff: float = 5.0 + + class ModelSection(ConfigSection): + num_interactions: int = 2 + radial: RadialSection = RadialSection() + + class TrainConfig(ReforgeBaseConfig): + seed: int = 1 + model: ModelSection = ModelSection() + +Here `model` and `radial` are sections. In a TOML file a section is a table +(`[model.radial]`), in YAML/JSON a nested mapping, and on the command line a +dotted prefix (`--model.radial.cutoff 5.0`). Values come from three layers, +lowest precedence first: + + schema defaults < one config file (.toml/.yaml/.yml/.json) < dotted CLI overrides + +Nothing else feeds a config: no environment variables, no dotenv files, so a +run is reproducible from its file and its command line alone. + +Unknown keys are hard errors. A CLI key is checked against the schema's +dotted paths before anything is built; a key in the file, or inside a +JSON-valued override, is caught by pydantic (`extra="forbid"` on every level +of the tree). Either way the message names the key by its dotted path and, +when there is one, the nearest valid neighbour. + +Field types are restricted to what survives a JSON round trip unchanged, so +that the resolved export is a fixed point: `set` and `frozenset` fields are +rejected when the schema class is defined, because their element order is not +stable across interpreter runs. Use a list. +""" + +from __future__ import annotations + +import difflib +import json +from collections.abc import Iterator, Sequence +from pathlib import Path +from types import UnionType +from typing import TYPE_CHECKING, Any, Union, get_args, get_origin + +import tomli +import yaml +from pydantic import BaseModel, ConfigDict, ValidationError + +if TYPE_CHECKING: + from typing_extensions import Self + +__all__ = ["ConfigError", "ConfigSection", "ReforgeBaseConfig", "read_config_file"] + +#: Config file extensions this module reads, keyed to their parsers. +_FILE_PARSERS = { + ".toml": tomli.loads, + ".yaml": yaml.safe_load, + ".yml": yaml.safe_load, + ".json": json.loads, +} + + +def _sections_in( + annotation: Any, inside: bool = False +) -> Iterator[tuple[type[BaseModel], bool]]: + """Every section class a field annotation can hold (`Radial`, `Radial | None`, + `list[Radial]`), with whether it sits inside a dict/list/tuple, where an + error location has a key or index before the section's own field names.""" + origin = get_origin(annotation) # `list` for `list[X]`; None for a plain class + if origin is None: + if isinstance(annotation, type) and issubclass(annotation, BaseModel): + yield annotation, inside + elif origin in (Union, UnionType): + for arg in get_args(annotation): + yield from _sections_in(arg, inside) + elif origin in (dict, list, tuple): + for arg in get_args(annotation): + yield from _sections_in(arg, True) + + +def _section_of(annotation: Any) -> tuple[type[BaseModel] | None, bool]: + """The one section a field can hold, and whether it is inside a collection.""" + found = dict(_sections_in(annotation)) + return next(iter(found.items())) if found else (None, False) + + +def _contains_a_set(annotation: Any) -> bool: + """A bare `set`, a `set[X]`, or a set anywhere inside, e.g. `list[set[int]]`.""" + if annotation in (set, frozenset) or get_origin(annotation) in (set, frozenset): + return True + return any(_contains_a_set(arg) for arg in get_args(annotation)) + + +def _check_schema(model: type[BaseModel]) -> None: + """Fail at class definition for a field shape the contract cannot keep. + + Each rule protects one guarantee: no sets (order is not stable across + runs, so the export would not be a fixed point); no aliases or computed + fields (the export would not validate back); one section class per field + (a value must not become whichever alternative happens to accept it, and + the CLI needs one set of valid keys under a field); every section is a + `ConfigSection` (a plain `BaseModel` ignores unknown keys, so a typo + would vanish, and skips these checks). + """ + + def reject(name: str, reason: str) -> None: + raise TypeError(f"{model.__name__}.{name} {reason}") + + for name in model.model_computed_fields: + reject(name, "is a computed field; the export must validate back, so drop it") + for name, field in model.model_fields.items(): + if field.alias or field.validation_alias or field.serialization_alias: + reject(name, "has an alias; config keys are field names, so drop it") + if _contains_a_set(field.annotation): + reject( + name, + "is typed as a set; set order is not stable across runs. Use a list", + ) + sections = {section for section, _ in _sections_in(field.annotation)} + if len(sections) > 1: + reject( + name, + "is a union of sections; give each alternative its own optional " + "field, e.g. `huber: HuberLoss | None`", + ) + for section in sections: + if not issubclass(section, ConfigSection): + reject( + name, + f"holds {section.__name__}, which is not a ConfigSection; " + f"subclass it", + ) + + +class ConfigError(ValueError): + """A config file or override the schema rejects. + + Raised for a missing, unparsable or malformed file, an unknown key and an + override the CLI parser cannot make sense of. The message names the file + or the offending key by its dotted path, and suggests the nearest valid + key when there is a close match. + """ + + +class ConfigSection(BaseModel): + """A nested section of a configuration: a table in the file, a dotted + prefix on the command line. Unknown keys are errors here too.""" + + model_config = ConfigDict(extra="forbid") + + @classmethod + def __pydantic_init_subclass__(cls, **kwargs: Any) -> None: + super().__pydantic_init_subclass__(**kwargs) + _check_schema(cls) + + +class ReforgeBaseConfig(ConfigSection): + """Root of a configuration tree. Subclass it; nest `ConfigSection`s in it. + + `load()` reads a file and applies overrides. Constructing the class + directly behaves like a plain Pydantic model. + """ + + @classmethod + def load( + cls, + config_file: str | Path | None = None, + cli_overrides: Sequence[str] = (), + ) -> Self: + """Build the config from defaults, then the file, then the overrides. + + `cli_overrides` is the argument list after the program name, e.g. + `["--model.num_interactions", "3", "--seed=7"]`: a dotted path names + a field at any depth. A value starting with `[` or `{`, or the word + `null`, is JSON, so a whole section, a list or a dict can be given; + any other value is a string pydantic converts to the field's type. + An override merges into the file, and into earlier overrides, like a + section does: a dict-valued field gains or replaces entries, so an + entry cannot be removed from the command line; a list-valued field + is replaced whole. + + Raises `ConfigError` for an unknown key, an unreadable file or an + unparsable override, and pydantic's `ValidationError` for a value of + the wrong type. + """ + values: dict[str, Any] = {} + if config_file is not None: + values = read_config_file(config_file) + values = _apply_overrides(cls, values, cli_overrides) + try: + return cls.model_validate(values) + except ValidationError as error: + unknown = _unknown_key_messages(cls, error) + if not unknown: + raise + raise ConfigError("\n".join(unknown)) from error + + def to_resolved_dict(self) -> dict[str, Any]: + """Every field, defaults included, as JSON-native values, in schema + order. Loading the result back and resolving again gives the same + dict. (TOML has no null, so a `None` can only go out as YAML or JSON.)""" + return self.model_dump(mode="json") + + def to_user_dict(self) -> dict[str, Any]: + """Only the fields the file and the overrides set, as JSON-native + values: what the user wrote, for the model metadata.""" + return self.model_dump(mode="json", exclude_unset=True) + + +def read_config_file(path: str | Path) -> dict[str, Any]: + """Parse one config file, choosing the parser by extension. + + An empty TOML or YAML file is an empty config. Raises `ConfigError` for a missing + file, an unknown extension, a file its parser rejects, or a file whose + top level is not a table. + """ + path = Path(path) + parser = _FILE_PARSERS.get(path.suffix.lower()) + if parser is None: + raise ConfigError( + f"cannot read config file {path}: unknown extension {path.suffix!r}; " + f"expected one of {', '.join(_FILE_PARSERS)}" + ) + try: + values = parser(path.read_text(encoding="utf-8")) + except OSError as error: + raise ConfigError( + f"cannot read config file {path}: {error.strerror}" + ) from error + except (ValueError, yaml.YAMLError) as error: # tomli/json errors are ValueErrors + raise ConfigError(f"cannot parse config file {path}: {error}") from error + if values is None: + return {} + # TOML always yields a table, but a YAML or JSON file can hold a list or a + # scalar, which would crash the merge with the overrides instead of + # naming the file. + if not isinstance(values, dict): + raise ConfigError( + f"config file {path} must hold a table of keys at the top level, " + f"not a {type(values).__name__}" + ) + return values + + +def _deep_update(base: dict[str, Any], update: dict[str, Any]) -> dict[str, Any]: + """`base` overlaid with `update`, recursing where both hold a dict.""" + merged = dict(base) + for key, value in update.items(): + if isinstance(value, dict) and isinstance(merged.get(key), dict): + merged[key] = _deep_update(merged[key], value) + else: + merged[key] = value + return merged + + +def _apply_overrides( + model: type[BaseModel], values: dict[str, Any], cli_overrides: Sequence[str] +) -> dict[str, Any]: + """`values` with the `--a.b value` and `--a.b=value` pairs merged in, in order.""" + valid = list(_dotted_paths(model)) + valid_set = set(valid) + unknown: list[str] = [] + tokens = iter(cli_overrides) + + for token in tokens: + option = token.removeprefix("--") + if "=" in option: + name, value = option.split("=", 1) + else: + name, value = option, None + if not token.startswith("--") or not name: + raise ConfigError( + f"unknown config option {token!r}; overrides are written --key value" + ) + if value is None: + value = next(tokens, None) + if value is None: + raise ConfigError(f"override --{name} is missing its value") + + if name not in valid_set: + unknown.append(_unknown_key_message(name, valid)) + continue + + if value == "null" or value.startswith(("[", "{")): + try: + value = json.loads(value) + except ValueError as error: + raise ConfigError( + f"override --{name} is not valid JSON: {error}" + ) from error + + # Nested under its path, an override merges by the file's rule. + override: Any = value + for section in reversed(name.split(".")): + override = {section: override} + values = _deep_update(values, override) + + if unknown: + raise ConfigError("\n".join(unknown)) + + return values + + +def _dotted_paths(model: type[BaseModel], prefix: str = "") -> Iterator[str]: + """Every field of the tree as a dotted path, sections included. + + A section inside a dict or list is not descended into: the CLI addresses + such a field only as a whole, with a JSON value. + """ + for name, field in model.model_fields.items(): + path = f"{prefix}{name}" + yield path + section, inside_collection = _section_of(field.annotation) + if section is not None and not inside_collection: + yield from _dotted_paths(section, f"{path}.") + + +def _unknown_key_message(key: str, candidates: Sequence[str]) -> str: + message = f"unknown config key {key!r}" + closest = difflib.get_close_matches(key, candidates, n=1) + if closest: + message += f"; did you mean {closest[0]!r}?" + return message + + +def _unknown_key_messages(model: type[BaseModel], error: ValidationError) -> list[str]: + """Pydantic's unknown-key errors as messages that name the full dotted + path and the closest valid key at that level. + + The error location is walked against the schema to find the section whose + fields are the candidates: a field name moves into its section, a dict key + or list index stays in it, and the tag pydantic inserts for a + `Section | scalar` field (the class name) is dropped from the path. + """ + messages = [] + for item in error.errors(): + # pydantic's error code for a key that matches no field (extra="forbid"). + # Every other code, e.g. a wrong type, is left for `load()` to re-raise. + if item["type"] != "extra_forbidden": + continue + *location, key = (str(part) for part in item["loc"]) + section: type[BaseModel] | None = model + inside_collection = False + names = [] + for part in location: + if section is not None and part == section.__name__: + continue + names.append(part) + if inside_collection: + inside_collection = False + elif section is not None and part in section.model_fields: + section, inside_collection = _section_of( + section.model_fields[part].annotation + ) + else: + section = None + candidates = list(section.model_fields) if section is not None else [] + prefix = "".join(f"{name}." for name in names) + messages.append( + _unknown_key_message(f"{prefix}{key}", [f"{prefix}{c}" for c in candidates]) + ) + return messages diff --git a/packages/mace-core/src/mace_core/metadata.py b/packages/mace-core/src/mace_core/metadata.py new file mode 100644 index 000000000..bc016565c --- /dev/null +++ b/packages/mace-core/src/mace_core/metadata.py @@ -0,0 +1,237 @@ +"""The record every v1 model carries about how it was made. + +`ModelMetadata` is stored alongside the weights of every trained model, not +only foundation models. It is plain data: a Pydantic tree that serialises to +JSON with `to_json()` and comes back, without loss, through `from_json()`. + +The record is versioned. `SCHEMA_VERSION` is bumped whenever a field is +added, removed or changes meaning, and `from_json()` refuses a record written +under a version this code does not know, so a newer checkpoint fails loudly +at load time instead of being read with the wrong meanings. +""" + +from __future__ import annotations + +import json +from collections.abc import Iterable +from typing import Any, Final, Literal + +from pydantic import BaseModel, ConfigDict, Field + +from mace_core.config import ReforgeBaseConfig + +__all__ = [ + "SCHEMA_VERSION", + "Citation", + "ConfigRecord", + "DataSourceSummary", + "DataSummary", + "E0Details", + "MetadataSchemaError", + "ModelMetadata", + "ParentModel", + "Provenance", + "format_citations", +] + +#: The schema version this module writes and the only one it reads. Bump it +#: together with the `Literal` on `ModelMetadata.schema_version`. +SCHEMA_VERSION: Final = 1 + + +class MetadataSchemaError(ValueError): + """The metadata was written under a schema version this code cannot read.""" + + +class _Record(BaseModel): + """Common ground: unknown keys are errors, so a typo cannot be stored.""" + + # inf/nan are written as JSON constants (Infinity, NaN) rather than + # pydantic's default null, which would turn a value into a different one. + model_config = ConfigDict(extra="forbid", ser_json_inf_nan="constants") + + +class ConfigRecord(_Record): + """The training configuration, as written and as resolved. + + Both are the JSON-native dicts a `ReforgeBaseConfig` exports: `user` is + `to_user_dict()`, the keys the config file and the command line set, and + `resolved` is `to_resolved_dict()`, every key with defaults filled in. + Build it with `from_config()` so the two cannot be mixed up. + """ + + user: dict[str, Any] = Field(default_factory=dict) + resolved: dict[str, Any] = Field(default_factory=dict) + + @classmethod + def from_config(cls, config: ReforgeBaseConfig) -> ConfigRecord: + return cls(user=config.to_user_dict(), resolved=config.to_resolved_dict()) + + +class Provenance(_Record): + """Which code produced the model.""" + + #: Version of the `mace-core` distribution, as `mace_core.__version__` reports it. + code_version: str + #: Full hash of the commit the code was run from; None when not in a checkout. + git_commit: str | None = None + + +class DataSourceSummary(_Record): + """Automated summary of one data source. + + Reference-quantity keys name the method that produced the reference as a + prefix on the quantity: `pbe_energy`, `pbe_forces`, `r2scan_energy`. + `reference_keys` lists the keys this source was fitted to, under that + convention, so a reader can tell which level of theory each head + reproduces. Which heads a source feeds is in the resolved config. + """ + + #: The data source's name in the config. + name: str + num_configurations: int | None = None + num_atoms: int | None = None + #: Chemical symbols of every element present. + elements: list[str] = Field(default_factory=list) + reference_keys: list[str] = Field(default_factory=list) + + +class DataSummary(_Record): + """One summary per data source; totals are sums over them, not stored.""" + + sources: list[DataSourceSummary] = Field(default_factory=list) + + +class E0Details(_Record): + """How one head's per-element reference energies (the E0s) were obtained. + + `values` maps chemical symbol to E0 in the model's energy unit. Symbols + rather than atomic numbers, because JSON keys are strings and an integer + key would not survive the round trip. + """ + + #: "explicit" when the E0s were given; "estimated" when fitted from the data. + source: Literal["explicit", "estimated"] + #: The estimation method (e.g. "average", "least_squares"); None when explicit. + method: str | None = None + #: Parameters of the method or reference, e.g. which data the fit used. + parameters: dict[str, Any] = Field(default_factory=dict) + values: dict[str, float] = Field(default_factory=dict) + + +class Citation(_Record): + """One work users of the model are asked to cite.""" + + title: str + authors: list[str] = Field(default_factory=list) + venue: str | None = None + year: int | None = None + doi: str | None = None + url: str | None = None + + +class ModelMetadata(_Record): + """The mandatory per-model record. See the module docstring.""" + + #: Pinned to the version this code reads; a bump here is a schema change. + schema_version: Literal[1] = SCHEMA_VERSION + config: ConfigRecord + provenance: Provenance + data: DataSummary = Field(default_factory=DataSummary) + #: Keyed by head name, as in the config; every head has its own E0 table. + e0: dict[str, E0Details] = Field(default_factory=dict) + #: DOI of the model itself, not of the papers describing it. + doi: str | None = None + citations: list[Citation] = Field(default_factory=list) + notes: str = "" + #: The models this one was built from, each with its own record inside, so + #: the whole lineage travels with the model. + parents: list[ParentModel] = Field(default_factory=list) + + def to_json(self, indent: int | None = 2) -> str: + """Serialise; raises `MetadataSchemaError` if the text would not read + back to an equal record, so a lossy value (a tuple that comes back as + a list, a datetime that comes back as a string) is never stored.""" + text = self.model_dump_json(indent=indent) + if self.from_json(text) != self: + raise MetadataSchemaError( + "model metadata does not survive a JSON round trip; " + "a field holds a value JSON cannot represent" + ) + return text + + @classmethod + def from_json(cls, text: str) -> ModelMetadata: + """Parse a record written by `to_json()`. + + Raises `MetadataSchemaError` when the record carries a schema version + this code does not read, before any field is interpreted. + """ + try: + document = json.loads(text) + except ValueError as error: + raise MetadataSchemaError( + f"model metadata is not valid JSON: {error}" + ) from error + if not isinstance(document, dict): + raise MetadataSchemaError( + f"model metadata must be a JSON object, not {type(document).__name__}" + ) + version = document.get("schema_version") + if type(version) is not int: + raise MetadataSchemaError( + f"model metadata has schema_version {version!r}; expected the " + f"integer {SCHEMA_VERSION}" + ) + if version != SCHEMA_VERSION: + hint = ( + "it was written by a newer mace-core; upgrade to read it" + if version > SCHEMA_VERSION + else "no migration exists for it" + ) + raise MetadataSchemaError( + f"model metadata has schema_version {version}, but this " + f"mace-core reads schema_version {SCHEMA_VERSION}: {hint}" + ) + return cls.model_validate_json(text) + + +class ParentModel(_Record): + """A model this one was built from.""" + + #: What the parent contributed: its weights as the starting point (fine-tuning + #: or continued training), or its predictions as distillation targets. + role: Literal["initial_weights", "teacher"] + #: How the config named it: a path or a registry name such as "mace-mp-0b3". + name: str + #: The parent's own record; None for a legacy checkpoint that carries none. + metadata: ModelMetadata | None = None + + +ModelMetadata.model_rebuild() # `parents` refers to the class defined after it + + +def format_citations(citations: Iterable[Citation]) -> str: + """Render citations as a numbered, printable block; empty for none. + + One line per citation: authors, title, venue and year, then the DOI or + URL. Fields that are unset are left out rather than printed as None. + """ + lines = [] + for number, citation in enumerate(citations, start=1): + parts = [] + if citation.authors: + parts.append(", ".join(citation.authors)) + parts.append(citation.title) + if citation.venue and citation.year: + parts.append(f"{citation.venue} ({citation.year})") + elif citation.venue: + parts.append(citation.venue) + elif citation.year: + parts.append(str(citation.year)) + if citation.doi: + parts.append(f"https://doi.org/{citation.doi}") + elif citation.url: + parts.append(citation.url) + lines.append(f"[{number}] " + ". ".join(parts)) + return "\n".join(lines) diff --git a/packages/mace-core/tests/test_mace_core_config.py b/packages/mace-core/tests/test_mace_core_config.py new file mode 100644 index 000000000..70c02c124 --- /dev/null +++ b/packages/mace-core/tests/test_mace_core_config.py @@ -0,0 +1,475 @@ +"""`ReforgeBaseConfig`: file formats, precedence, dotted overrides, unknown keys, +and the resolved export's fixed point.""" + +import json +import subprocess +import sys +from typing import Annotated + +import pytest +import yaml +from mace_core.config import ConfigError, ConfigSection, ReforgeBaseConfig +from pydantic import BaseModel, Field, ValidationError, computed_field + +# --------------------------------------------------------------------------- +# The demo schema: two levels of nesting, a list, an optional, a Literal. + + +class RadialSection(ConfigSection): + num_bessel: int = 8 + cutoff: float = 5.0 + + +class ModelSection(ConfigSection): + num_interactions: int = 2 + hidden_irreps: str = "128x0e + 128x1o" + radial: RadialSection = RadialSection() + + +class DataSection(ConfigSection): + train_file: str | None = None + valid_fraction: float = 0.1 + energy_key: str = "REF_energy" + heads: list[str] = Field(default_factory=lambda: ["default"]) + + +class StageTwoSection(ConfigSection): + start_epoch: int = 100 + energy_weight: float = 1000.0 + + +class DemoConfig(ReforgeBaseConfig): + name: str = "mace" + seed: int = 123 + default_dtype: str = "float64" + model: ModelSection = ModelSection() + data: DataSection = DataSection() + #: An optional section: absent unless the file or the CLI opens it. + stage_two: StageTwoSection | None = None + + +#: One config, as a dict. Each format test writes it out and loads it back. +FILE_VALUES = { + "name": "water", + "seed": 7, + "model": {"num_interactions": 4, "radial": {"cutoff": 4.5}}, + "data": {"train_file": "train.xyz", "heads": ["pbe", "r2scan"]}, +} + + +def to_toml(values, prefix=""): + """Enough TOML for a None-free config: scalars and lists share JSON's + literal syntax, nested dicts become `[a.b]` tables after the scalars.""" + lines = [ + f"{k} = {json.dumps(v)}" for k, v in values.items() if not isinstance(v, dict) + ] + for key, value in values.items(): + if isinstance(value, dict): + lines += [f"\n[{prefix}{key}]", to_toml(value, f"{prefix}{key}.")] + return "\n".join(lines) + + +def dump(values, extension): + if extension == ".toml": + return to_toml(values) + if extension == ".json": + return json.dumps(values) + return yaml.safe_dump(values) + + +def write_config(tmp_path, extension, values=FILE_VALUES): + path = tmp_path / f"config{extension}" + path.write_text(dump(values, extension), encoding="utf-8") + return path + + +# --------------------------------------------------------------------------- +# File loading + + +@pytest.mark.parametrize("extension", [".toml", ".yaml", ".yml", ".json"]) +def test_same_config_loads_identically_from_every_format(tmp_path, extension): + config = DemoConfig.load(write_config(tmp_path, extension)) + assert config == DemoConfig.model_validate(FILE_VALUES) + # The file set two fields at depth two; the sibling kept its default. + assert config.model.radial.cutoff == 4.5 + assert config.model.radial.num_bessel == 8 + + +def test_unknown_extension_is_an_error(tmp_path): + path = tmp_path / "config.ini" + path.write_text("seed = 1", encoding="utf-8") + with pytest.raises(ConfigError, match=r"unknown extension '\.ini'"): + DemoConfig.load(path) + + +def test_empty_file_is_all_defaults(tmp_path): + path = tmp_path / "empty.yaml" + path.write_text("", encoding="utf-8") + assert DemoConfig.load(path) == DemoConfig() + + +def test_file_must_be_a_table_at_the_top(tmp_path): + path = tmp_path / "list.json" + path.write_text("[1, 2]", encoding="utf-8") + with pytest.raises(ConfigError, match="table of keys at the top level"): + DemoConfig.load(path) + + +def test_missing_file_is_a_config_error(tmp_path): + with pytest.raises(ConfigError, match=r"cannot read config file .*nope\.yaml"): + DemoConfig.load(tmp_path / "nope.yaml") + + +@pytest.mark.parametrize( + ("extension", "text"), + [(".toml", "seed = \n"), (".yaml", "seed: [1\n"), (".json", "{")], +) +def test_malformed_file_is_a_config_error(tmp_path, extension, text): + path = tmp_path / f"broken{extension}" + path.write_text(text, encoding="utf-8") + with pytest.raises(ConfigError, match=r"cannot parse config file .*broken"): + DemoConfig.load(path) + + +# --------------------------------------------------------------------------- +# Precedence: defaults < file < CLI. The legacy behaviour this pins is +# tests/unit/test_arg_parser.py::test_cli_flag_overrides_yaml_config. + + +def test_no_inputs_gives_the_defaults(): + config = DemoConfig.load() + assert config == DemoConfig() + assert config.model.num_interactions == 2 + + +def test_file_overrides_defaults(tmp_path): + config = DemoConfig.load(write_config(tmp_path, ".yaml")) + assert config.model.num_interactions == 4 # from the file + assert config.default_dtype == "float64" # untouched default + + +def test_cli_overrides_file_which_overrides_defaults(tmp_path): + config = DemoConfig.load( + write_config(tmp_path, ".toml"), ["--model.num_interactions", "3"] + ) + assert config.model.num_interactions == 3 # CLI beats the file's 4 + assert config.model.radial.cutoff == 4.5 # the file's other values survive + assert config.seed == 7 + assert config.model.radial.num_bessel == 8 # defaults fill the rest + assert config.default_dtype == "float64" + + +# --------------------------------------------------------------------------- +# Dotted CLI overrides + + +def test_dotted_override_reaches_a_two_level_nested_field(): + config = DemoConfig.load(cli_overrides=["--model.radial.cutoff", "6.0"]) + assert config.model.radial.cutoff == 6.0 + assert config.model.radial.num_bessel == 8 + + +def test_dotted_override_opens_an_optional_section(): + config = DemoConfig.load(cli_overrides=["--stage_two.start_epoch", "50"]) + assert config.stage_two == StageTwoSection(start_epoch=50) + assert DemoConfig.load().stage_two is None + + +def test_override_forms_and_types(): + config = DemoConfig.load( + cli_overrides=[ + "--seed=9", + "--data.train_file", + "null", + "--data.heads", + '["a", "b"]', + ] + ) + assert config.seed == 9 + assert config.data.train_file is None + assert config.data.heads == ["a", "b"] + + +def test_value_of_the_wrong_type_is_a_validation_error(tmp_path): + with pytest.raises(ValidationError, match="seed"): + DemoConfig.load(cli_overrides=["--seed", "seven"]) + with pytest.raises(ValidationError, match="seed"): + DemoConfig.load(write_config(tmp_path, ".yaml", {"seed": "seven"})) + + +def test_override_missing_its_value_is_a_config_error(): + with pytest.raises(ConfigError, match="override --seed is missing its value"): + DemoConfig.load(cli_overrides=["--seed"]) + + +def test_override_that_is_not_valid_json_is_a_config_error(): + with pytest.raises(ConfigError, match="override --model is not valid JSON"): + DemoConfig.load(cli_overrides=["--model", "{oops"]) + + +def test_value_starting_with_dashes_works_in_both_forms(): + assert DemoConfig.load(cli_overrides=["--name=--odd"]).name == "--odd" + assert DemoConfig.load(cli_overrides=["--name", "--odd"]).name == "--odd" + + +def test_dict_valued_field_takes_json_and_is_not_dotted_into(): + class Sources(ReforgeBaseConfig): + by_name: dict[str, RadialSection] = Field(default_factory=dict) + + config = Sources.load(cli_overrides=["--by_name", '{"pbe": {"cutoff": 4.0}}']) + assert config.by_name == {"pbe": RadialSection(cutoff=4.0)} + with pytest.raises(ConfigError, match=r"unknown config key 'by_name\.pbe\.cutoff'"): + Sources.load(cli_overrides=["--by_name.pbe.cutoff", "4.0"]) + # Inside an entry, the neighbour is still found: the key passes through. + with pytest.raises( + ConfigError, + match=r"'by_name\.pbe\.cutof'; did you mean 'by_name\.pbe\.cutoff'\?", + ): + Sources.load(cli_overrides=["--by_name", '{"pbe": {"cutof": 4.0}}']) + + +def test_dict_override_merges_entries_but_list_override_replaces(tmp_path): + class Sources(ReforgeBaseConfig): + by_name: dict[str, RadialSection] = Field(default_factory=dict) + heads: list[str] = Field(default_factory=list) + + path = write_config( + tmp_path, ".yaml", {"by_name": {"pbe": {"cutoff": 4.0}}, "heads": ["a", "b"]} + ) + config = Sources.load( + path, ["--by_name", '{"r2scan": {"cutoff": 6.0}}', "--heads", '["c"]'] + ) + assert set(config.by_name) == {"pbe", "r2scan"} + assert config.heads == ["c"] + + +def test_overrides_apply_in_order_on_top_of_the_file(tmp_path): + # Closing a section with null and reopening it drops what the file set + # in it; a dotted value followed by the whole section keeps both. + config = DemoConfig.load( + write_config(tmp_path, ".yaml", {"stage_two": {"energy_weight": 5.0}}), + cli_overrides=[ + "--stage_two", + "null", + "--stage_two.start_epoch", + "5", + "--model.radial.cutoff", + "4", + "--model", + '{"num_interactions": 3}', + ], + ) + assert config.stage_two == StageTwoSection(start_epoch=5) + assert (config.model.num_interactions, config.model.radial.cutoff) == (3, 4.0) + + +# --------------------------------------------------------------------------- +# Unknown keys name the key and its nearest neighbour, in files and on the CLI. + + +def test_unknown_top_level_key_in_file(tmp_path): + path = tmp_path / "typo.yaml" + path.write_text("sead: 1\n", encoding="utf-8") + with pytest.raises(ConfigError, match=r"'sead'; did you mean 'seed'\?"): + DemoConfig.load(path) + + +def test_unknown_nested_key_in_file_names_the_dotted_path(tmp_path): + path = tmp_path / "typo.json" + path.write_text(json.dumps({"model": {"radial": {"cutof": 4.0}}}), encoding="utf-8") + with pytest.raises( + ConfigError, + match=r"'model\.radial\.cutof'; did you mean 'model\.radial\.cutoff'\?", + ): + DemoConfig.load(path) + + +def test_every_unknown_key_is_reported_at_once(tmp_path): + path = tmp_path / "typos.yaml" + path.write_text("sead: 1\nmodel:\n num_interaction: 3\n", encoding="utf-8") + with pytest.raises(ConfigError) as excinfo: + DemoConfig.load(path) + assert "'sead'" in str(excinfo.value) + assert "'model.num_interaction'" in str(excinfo.value) + + +def test_unknown_key_without_a_close_neighbour_still_names_it(tmp_path): + path = tmp_path / "far.yaml" + path.write_text("zzzzzz: 1\n", encoding="utf-8") + with pytest.raises(ConfigError, match=r"unknown config key 'zzzzzz'$"): + DemoConfig.load(path) + + +def test_unknown_dotted_override_names_the_neighbour(): + with pytest.raises( + ConfigError, + match=r"'model\.num_interaction'; did you mean 'model\.num_interactions'\?", + ): + DemoConfig.load(cli_overrides=["--model.num_interaction", "3"]) + + +def test_unknown_key_inside_an_optional_section(): + with pytest.raises( + ConfigError, + match=r"'stage_two\.start'; did you mean 'stage_two\.start_epoch'\?", + ): + DemoConfig.load(cli_overrides=["--stage_two.start", "50"]) + + +def test_unknown_key_under_a_section_or_scalar_field_drops_the_tag(): + class SectionOrInt(ReforgeBaseConfig): + radial: RadialSection | int = 3 + + # pydantic tags the location with the member's class name; not a key. + with pytest.raises( + ConfigError, match=r"'radial\.cutof'; did you mean 'radial\.cutoff'\?" + ): + SectionOrInt.load(cli_overrides=["--radial", '{"cutof": 4.0}']) + + +def test_help_flag_is_an_error_not_an_exit(): + with pytest.raises(ConfigError, match=r"unknown config option '-h'"): + DemoConfig.load(cli_overrides=["-h"]) + + +def test_abbreviated_option_is_unknown_not_expanded(): + with pytest.raises(ConfigError, match=r"key 'se'; did you mean 'seed'"): + DemoConfig.load(cli_overrides=["--se", "3"]) + + +def test_empty_inline_value_does_not_hide_the_next_option(): + with pytest.raises(ConfigError, match=r"'sead'; did you mean 'seed'\?"): + DemoConfig.load(cli_overrides=["--name=", "--sead", "5"]) + + +def test_direct_construction_rejects_unknown_keys_too(): + with pytest.raises(ValidationError, match="extra_forbidden"): + DemoConfig(model={"num_interaction": 3}) + + +# --------------------------------------------------------------------------- +# Resolved export + + +def test_resolved_dict_has_every_default_in_declaration_order(tmp_path): + # The file lists keys in the reverse of the schema's order. + path = tmp_path / "reversed.yaml" + path.write_text("seed: 1\nname: x\n", encoding="utf-8") + resolved = DemoConfig.load(path).to_resolved_dict() + assert list(resolved) == [ + "name", + "seed", + "default_dtype", + "model", + "data", + "stage_two", + ] + assert resolved["stage_two"] is None + assert list(resolved["model"]) == ["num_interactions", "hidden_irreps", "radial"] + assert resolved["model"]["radial"] == {"num_bessel": 8, "cutoff": 5.0} + assert resolved["data"]["train_file"] is None + + +def assert_fixed_point(tmp_path, first, extension): + written = tmp_path / f"resolved{extension}" + written.write_text(dump(first, extension), encoding="utf-8") + second = DemoConfig.load(written).to_resolved_dict() + assert second == first + assert json.dumps(second) == json.dumps(first) # order included + + +@pytest.mark.parametrize("extension", [".yaml", ".json"]) +def test_file_to_resolved_to_file_to_resolved_is_a_fixed_point(tmp_path, extension): + first = DemoConfig.load( + write_config(tmp_path, ".toml"), ["--model.num_interactions", "3"] + ).to_resolved_dict() + assert first["stage_two"] is None # a None is part of what has to survive + assert_fixed_point(tmp_path, first, extension) + + +def test_fixed_point_holds_through_toml_when_nothing_is_none(tmp_path): + # TOML has no null, so the optional section is opened and the optional + # file name set; the resolved dict then goes through all three formats. + first = DemoConfig.load( + write_config(tmp_path, ".yaml"), ["--stage_two.start_epoch", "50"] + ).to_resolved_dict() + assert "null" not in json.dumps(first) + for extension in (".toml", ".yaml", ".json"): + assert_fixed_point(tmp_path, first, extension) + + +class LenientSection(BaseModel): + cutoff: float = 5.0 + + +def test_field_shapes_the_contract_cannot_keep_are_rejected_at_class_definition(): + # Each shape would break a guarantee: set order varies with the hash + # seed; aliases and computed fields do not validate back; a union of + # sections would let a value pick its section; a lenient section would + # swallow typos. + shapes = { + r"tags is typed as a set.*Use a list": ("tags", list[set[str]]), + r"num has an alias": ("num", Annotated[int, Field(alias="n")]), + r"either is a union of sections": ("either", RadialSection | StageTwoSection), + r"radial holds LenientSection, which is not a ConfigSection": ( + "radial", + LenientSection | None, + ), + } + for message, (name, annotation) in shapes.items(): + with pytest.raises(TypeError, match=message): + type("Bad", (ConfigSection,), {"__annotations__": {name: annotation}}) + + with pytest.raises(TypeError, match=r"double is a computed field"): + + class Computed(ConfigSection): + seed: int = 1 + + @computed_field + def double(self) -> int: + return 2 * self.seed + + +def test_user_dict_holds_only_what_was_set(tmp_path): + config = DemoConfig.load( + write_config(tmp_path, ".json"), ["--model.num_interactions", "3"] + ) + assert config.to_user_dict() == { + "name": "water", + "seed": 7, + "model": {"num_interactions": 3, "radial": {"cutoff": 4.5}}, + "data": {"train_file": "train.xyz", "heads": ["pbe", "r2scan"]}, + } + + +# --------------------------------------------------------------------------- +# Nothing but the file and the CLI feeds a config. + + +def test_environment_variables_are_ignored(monkeypatch): + monkeypatch.setenv("NAME", "from-the-environment") + monkeypatch.setenv("SEED", "99") + config = DemoConfig.load() + assert config.name == "mace" + assert config.seed == 123 + + +@pytest.mark.parametrize("key", ["SEED", "Seed", "_env_file", "_cli_parse_args"]) +def test_root_keys_are_validated_like_any_section(tmp_path, key): + # Neither case variants nor BaseSettings-style private constructor + # options are special at the top level: unknown is unknown. + path = tmp_path / "root.json" + path.write_text(json.dumps({key: 1}), encoding="utf-8") + with pytest.raises(ConfigError, match=f"unknown config key '{key}'"): + DemoConfig.load(path) + + +def test_config_module_imports_neither_torch_nor_jax(): + """In a fresh interpreter, so another test's imports cannot mask a leak.""" + code = ( + "import sys, mace_core.config; " + "leaked = {'torch', 'jax', 'e3nn'} & set(sys.modules); " + "assert not leaked, leaked" + ) + subprocess.run([sys.executable, "-c", code], check=True) diff --git a/packages/mace-core/tests/test_mace_core_metadata.py b/packages/mace-core/tests/test_mace_core_metadata.py new file mode 100644 index 000000000..b560b529c --- /dev/null +++ b/packages/mace-core/tests/test_mace_core_metadata.py @@ -0,0 +1,234 @@ +"""`ModelMetadata`: JSON round trip, schema versioning, citation rendering.""" + +import json +import subprocess +import sys + +import pytest +from mace_core.config import ConfigSection, ReforgeBaseConfig +from mace_core.metadata import ( + SCHEMA_VERSION, + Citation, + ConfigRecord, + DataSourceSummary, + DataSummary, + E0Details, + MetadataSchemaError, + ModelMetadata, + ParentModel, + Provenance, + format_citations, +) +from pydantic import ValidationError + +MACE_PAPER = Citation( + title="MACE: Higher Order Equivariant Message Passing Neural Networks " + "for Fast and Accurate Force Fields", + authors=["I. Batatia", "D. P. Kovacs", "G. N. C. Simm", "C. Ortner", "G. Csanyi"], + venue="Advances in Neural Information Processing Systems", + year=2022, + url="https://arxiv.org/abs/2206.07697", +) + + +def full_record() -> ModelMetadata: + """Every field set, so the round trip is tested on all of them.""" + return ModelMetadata( + config=ConfigRecord( + user={"model": {"num_interactions": 3}}, + resolved={"name": "mace", "model": {"num_interactions": 3, "cutoff": 5.0}}, + ), + provenance=Provenance(code_version="1.0.0", git_commit="a" * 40), + data=DataSummary( + sources=[ + DataSourceSummary( + name="water", + num_configurations=1200, + num_atoms=64_000, + elements=["H", "O"], + reference_keys=["pbe_energy", "pbe_forces"], + ), + DataSourceSummary(name="ice", elements=["H", "O"]), + ] + ), + e0={ + "pbe": E0Details( + source="estimated", + method="least_squares", + parameters={"reference_key": "pbe_energy"}, + values={"H": -13.6, "O": -430.2}, + ), + "r2scan": E0Details(source="explicit", values={"H": -13.7, "O": -431.0}), + }, + doi="10.5281/zenodo.0000000", + citations=[MACE_PAPER, Citation(title="A dataset paper", doi="10.1000/xyz")], + notes="Trained for the round-trip test.", + ) + + +# --------------------------------------------------------------------------- +# Round trip and schema version + + +def test_json_round_trip_is_lossless(): + record = full_record() + assert ModelMetadata.from_json(record.to_json()) == record + + +def test_minimal_record_round_trips_too(): + record = ModelMetadata( + config=ConfigRecord(), provenance=Provenance(code_version="0.0.0") + ) + assert ModelMetadata.from_json(record.to_json()) == record + assert record.e0 == {} + + +def test_config_and_provenance_are_mandatory(): + with pytest.raises(ValidationError, match="config"): + ModelMetadata.model_validate({"provenance": {"code_version": "0"}}) + + +def test_lossy_value_is_refused_rather_than_stored(): + record = full_record() + record.e0["pbe"].parameters["shape"] = (2, 3) # JSON brings it back as a list + with pytest.raises(MetadataSchemaError, match="does not survive a JSON round trip"): + record.to_json() + + +def test_lineage_round_trips_through_two_levels(): + foundation = ParentModel(role="initial_weights", name="mace-mp-0b3") # no record + distilled = full_record() + distilled.parents = [ + foundation, + ParentModel(role="teacher", name="teacher.model", metadata=full_record()), + ] + fine_tuned = full_record() + fine_tuned.parents = [ + ParentModel(role="initial_weights", name="distilled.model", metadata=distilled) + ] + back = ModelMetadata.from_json(fine_tuned.to_json()) + assert back == fine_tuned + assert back.parents[0].metadata is not None + grandparents = back.parents[0].metadata.parents + assert [p.role for p in grandparents] == ["initial_weights", "teacher"] + assert grandparents[0].metadata is None + + +def test_schema_version_is_written(): + assert json.loads(full_record().to_json())["schema_version"] == SCHEMA_VERSION + + +def test_future_schema_version_is_rejected_clearly(): + document = json.loads(full_record().to_json()) + document["schema_version"] = SCHEMA_VERSION + 1 + with pytest.raises(MetadataSchemaError) as excinfo: + ModelMetadata.from_json(json.dumps(document)) + message = str(excinfo.value) + assert f"schema_version {SCHEMA_VERSION + 1}" in message + assert f"reads schema_version {SCHEMA_VERSION}" in message + assert "upgrade" in message + + +@pytest.mark.parametrize( + ("mutate", "message"), + [ + ( + lambda d: d.pop("schema_version"), + "schema_version None; expected the integer 1", + ), + (lambda d: d.update(schema_version="1"), "schema_version '1'; expected"), + (lambda d: d.update(schema_version=1.0), "schema_version 1.0; expected"), + ], +) +def test_missing_or_non_integer_schema_version_is_rejected(mutate, message): + document = json.loads(full_record().to_json()) + mutate(document) + with pytest.raises(MetadataSchemaError, match=message): + ModelMetadata.from_json(json.dumps(document)) + + +@pytest.mark.parametrize( + ("text", "message"), + [("[1]", "must be a JSON object, not list"), ("{", "is not valid JSON")], +) +def test_non_record_json_is_rejected_with_context(text, message): + with pytest.raises(MetadataSchemaError, match=message): + ModelMetadata.from_json(text) + + +def test_infinity_survives_and_nan_is_refused(): + # pydantic's default writes inf/nan as null, which would silently turn an + # E0 into a different value. NaN is never equal to itself, so it cannot + # pass the round-trip check; an E0 or a config value that is NaN is a bug + # upstream, not something to store. + record = full_record() + record.e0["pbe"].values["H"] = float("inf") + back = ModelMetadata.from_json(record.to_json()) + assert back.e0["pbe"].values["H"] == float("inf") + record.config.resolved["cutoff"] = float("nan") + with pytest.raises(MetadataSchemaError, match="does not survive"): + record.to_json() + + +def test_schema_version_is_pinned_on_direct_validation_as_well(): + document = json.loads(full_record().to_json()) + document["schema_version"] = SCHEMA_VERSION + 1 + with pytest.raises(ValidationError, match="schema_version"): + ModelMetadata.model_validate(document) + + +def test_unknown_fields_are_rejected(): + with pytest.raises(ValidationError, match="extra_forbidden"): + ModelMetadata.model_validate( + {"config": {}, "provenance": {"code_version": "0"}, "note": "x"} + ) + + +def test_e0_source_is_one_of_two_values(): + with pytest.raises(ValidationError, match="source"): + E0Details.model_validate({"source": "guessed"}) + + +def test_config_record_is_built_from_a_config(): + class Section(ConfigSection): + cutoff: float = 5.0 + + class Config(ReforgeBaseConfig): + seed: int = 1 + model: Section = Section() + + record = ConfigRecord.from_config(Config.model_validate({"model": {"cutoff": 4.0}})) + assert record.user == {"model": {"cutoff": 4.0}} + assert record.resolved == {"seed": 1, "model": {"cutoff": 4.0}} + # The embedded form is the fixed point: resolving it again changes nothing. + assert Config.model_validate(record.resolved).to_resolved_dict() == record.resolved + + +# --------------------------------------------------------------------------- +# Citations + + +def test_citations_render_to_a_numbered_block(): + block = format_citations(full_record().citations) + assert block.splitlines() == [ + "[1] I. Batatia, D. P. Kovacs, G. N. C. Simm, C. Ortner, G. Csanyi. " + "MACE: Higher Order Equivariant Message Passing Neural Networks for " + "Fast and Accurate Force Fields. " + "Advances in Neural Information Processing Systems (2022). " + "https://arxiv.org/abs/2206.07697", + "[2] A dataset paper. https://doi.org/10.1000/xyz", + ] + + +def test_no_citations_render_to_nothing(): + assert format_citations([]) == "" + + +def test_metadata_module_imports_neither_torch_nor_jax(): + """In a fresh interpreter, so another test's imports cannot mask a leak.""" + code = ( + "import sys, mace_core.metadata; " + "leaked = {'torch', 'jax', 'e3nn'} & set(sys.modules); " + "assert not leaked, leaked" + ) + subprocess.run([sys.executable, "-c", code], check=True) From ed3913999cc1389e32e3f5954fd79e3a70ca3912 Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Thu, 17 Sep 2026 23:34:29 +0200 Subject: [PATCH 02/14] Record per head its E0s and the data sources it consumed (CORE-2, #1556) `ModelMetadata.heads` replaces the `e0` dict: a `HeadSummary` per head holds its `E0Details` and the names of the entries in `data.sources` it was fitted on, so the source-to-head wiring can be read from the record without the resolved config. A validator checks that every named source is summarised and that each source is summarised once, so a source shared by two heads is not counted twice. Docstrings: a source's `reference_keys` are the keys it provides, not what a head reproduces; `to_json` claims only that a text reading back to a different record is refused; a version mismatch inside an embedded parent record is pydantic's error, since a record only embeds parents it could read. --- packages/mace-core/README.md | 5 +- packages/mace-core/src/mace_core/metadata.py | 50 +++++++++++++++---- .../tests/test_mace_core_metadata.py | 40 +++++++++++---- 3 files changed, 71 insertions(+), 24 deletions(-) diff --git a/packages/mace-core/README.md b/packages/mace-core/README.md index f8f6999b2..83d2b013d 100644 --- a/packages/mace-core/README.md +++ b/packages/mace-core/README.md @@ -17,7 +17,8 @@ What is here so far: `to_resolved_dict()` for the fully defaulted, round-trippable export. - `mace_core.metadata` — `ModelMetadata`, the versioned record every trained model carries (config as written and as resolved, provenance, a summary per - data source, E0 details per head, DOI, citations, notes, and the records of - the models it was fine-tuned or distilled from), with a JSON round trip, + data source, per head its E0s and the sources it consumed, DOI, citations, + notes, and the records of the models it was built from), with a JSON round + trip, `ConfigRecord.from_config()` to embed a config in its fixed-point form, and `format_citations()`. diff --git a/packages/mace-core/src/mace_core/metadata.py b/packages/mace-core/src/mace_core/metadata.py index bc016565c..575732c99 100644 --- a/packages/mace-core/src/mace_core/metadata.py +++ b/packages/mace-core/src/mace_core/metadata.py @@ -16,7 +16,7 @@ from collections.abc import Iterable from typing import Any, Final, Literal -from pydantic import BaseModel, ConfigDict, Field +from pydantic import BaseModel, ConfigDict, Field, model_validator from mace_core.config import ReforgeBaseConfig @@ -27,6 +27,7 @@ "DataSourceSummary", "DataSummary", "E0Details", + "HeadSummary", "MetadataSchemaError", "ModelMetadata", "ParentModel", @@ -82,9 +83,8 @@ class DataSourceSummary(_Record): Reference-quantity keys name the method that produced the reference as a prefix on the quantity: `pbe_energy`, `pbe_forces`, `r2scan_energy`. - `reference_keys` lists the keys this source was fitted to, under that - convention, so a reader can tell which level of theory each head - reproduces. Which heads a source feeds is in the resolved config. + `reference_keys` lists the keys this source provides, under that + convention. The heads a source fed name it in `ModelMetadata.heads`. """ #: The data source's name in the config. @@ -97,7 +97,8 @@ class DataSourceSummary(_Record): class DataSummary(_Record): - """One summary per data source; totals are sums over them, not stored.""" + """One summary per data source, each once even when several heads share + it (`ModelMetadata` checks that); totals are sums over them, not stored.""" sources: list[DataSourceSummary] = Field(default_factory=list) @@ -119,6 +120,14 @@ class E0Details(_Record): values: dict[str, float] = Field(default_factory=dict) +class HeadSummary(_Record): + """What one head was fitted on: its E0s and the data sources it consumed.""" + + e0: E0Details + #: Names in `DataSummary.sources`; a source feeding two heads appears in both. + sources: list[str] = Field(default_factory=list) + + class Citation(_Record): """One work users of the model are asked to cite.""" @@ -138,8 +147,8 @@ class ModelMetadata(_Record): config: ConfigRecord provenance: Provenance data: DataSummary = Field(default_factory=DataSummary) - #: Keyed by head name, as in the config; every head has its own E0 table. - e0: dict[str, E0Details] = Field(default_factory=dict) + #: Keyed by head name, as in the config; a single-head model has one entry. + heads: dict[str, HeadSummary] = Field(default_factory=dict) #: DOI of the model itself, not of the papers describing it. doi: str | None = None citations: list[Citation] = Field(default_factory=list) @@ -148,10 +157,27 @@ class ModelMetadata(_Record): #: the whole lineage travels with the model. parents: list[ParentModel] = Field(default_factory=list) + @model_validator(mode="after") + def _heads_name_known_sources(self) -> ModelMetadata: + names = [source.name for source in self.data.sources] + if len(set(names)) != len(names): + raise ValueError( + f"data.sources names a source twice: {sorted(names)}; " + f"summarise each source once" + ) + for head, summary in self.heads.items(): + for name in summary.sources: + if name not in names: + raise ValueError( + f"heads.{head}.sources names {name!r}, which is not in " + f"data.sources; add its summary or drop the name" + ) + return self + def to_json(self, indent: int | None = 2) -> str: - """Serialise; raises `MetadataSchemaError` if the text would not read - back to an equal record, so a lossy value (a tuple that comes back as - a list, a datetime that comes back as a string) is never stored.""" + """Serialise; raises `MetadataSchemaError` if the text reads back to a + different record, so a lossy value (a tuple that comes back as a list, + a datetime that comes back as a string) is never stored.""" text = self.model_dump_json(indent=indent) if self.from_json(text) != self: raise MetadataSchemaError( @@ -165,7 +191,9 @@ def from_json(cls, text: str) -> ModelMetadata: """Parse a record written by `to_json()`. Raises `MetadataSchemaError` when the record carries a schema version - this code does not read, before any field is interpreted. + this code does not read, before any field is interpreted. Embedded + parent records are validated as fields, so a version mismatch inside + one is pydantic's error; a record only embeds parents it could read. """ try: document = json.loads(text) diff --git a/packages/mace-core/tests/test_mace_core_metadata.py b/packages/mace-core/tests/test_mace_core_metadata.py index b560b529c..4fa877f02 100644 --- a/packages/mace-core/tests/test_mace_core_metadata.py +++ b/packages/mace-core/tests/test_mace_core_metadata.py @@ -13,6 +13,7 @@ DataSourceSummary, DataSummary, E0Details, + HeadSummary, MetadataSchemaError, ModelMetadata, ParentModel, @@ -51,14 +52,20 @@ def full_record() -> ModelMetadata: DataSourceSummary(name="ice", elements=["H", "O"]), ] ), - e0={ - "pbe": E0Details( - source="estimated", - method="least_squares", - parameters={"reference_key": "pbe_energy"}, - values={"H": -13.6, "O": -430.2}, + heads={ + "pbe": HeadSummary( + e0=E0Details( + source="estimated", + method="least_squares", + parameters={"reference_key": "pbe_energy"}, + values={"H": -13.6, "O": -430.2}, + ), + sources=["water", "ice"], + ), + "r2scan": HeadSummary( + e0=E0Details(source="explicit", values={"H": -13.7, "O": -431.0}), + sources=["ice"], ), - "r2scan": E0Details(source="explicit", values={"H": -13.7, "O": -431.0}), }, doi="10.5281/zenodo.0000000", citations=[MACE_PAPER, Citation(title="A dataset paper", doi="10.1000/xyz")], @@ -80,7 +87,18 @@ def test_minimal_record_round_trips_too(): config=ConfigRecord(), provenance=Provenance(code_version="0.0.0") ) assert ModelMetadata.from_json(record.to_json()) == record - assert record.e0 == {} + assert record.heads == {} + + +def test_heads_must_name_summarised_sources(): + record = full_record() + record.heads["pbe"].sources.append("vapour") + with pytest.raises(ValidationError, match=r"heads\.pbe\.sources names 'vapour'"): + ModelMetadata.model_validate(record.model_dump()) + record = full_record() + record.data.sources.append(DataSourceSummary(name="ice")) + with pytest.raises(ValidationError, match="names a source twice"): + ModelMetadata.model_validate(record.model_dump()) def test_config_and_provenance_are_mandatory(): @@ -90,7 +108,7 @@ def test_config_and_provenance_are_mandatory(): def test_lossy_value_is_refused_rather_than_stored(): record = full_record() - record.e0["pbe"].parameters["shape"] = (2, 3) # JSON brings it back as a list + record.heads["pbe"].e0.parameters["shape"] = (2, 3) # JSON brings it back as a list with pytest.raises(MetadataSchemaError, match="does not survive a JSON round trip"): record.to_json() @@ -162,9 +180,9 @@ def test_infinity_survives_and_nan_is_refused(): # pass the round-trip check; an E0 or a config value that is NaN is a bug # upstream, not something to store. record = full_record() - record.e0["pbe"].values["H"] = float("inf") + record.heads["pbe"].e0.values["H"] = float("inf") back = ModelMetadata.from_json(record.to_json()) - assert back.e0["pbe"].values["H"] == float("inf") + assert back.heads["pbe"].e0.values["H"] == float("inf") record.config.resolved["cutoff"] = float("nan") with pytest.raises(MetadataSchemaError, match="does not survive"): record.to_json() From e5339d4f3a47f7840b087514b56b764857fc3678 Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Sun, 20 Sep 2026 21:53:50 +0200 Subject: [PATCH 03/14] Write a config section's kind as its key, the last kind written wins (CORE-2 follow-up, #1556) A config field may be a discriminated union of sections, "kinds": loss: Annotated[WeightedLoss | HuberLoss, Field(discriminator="kind")] = WeightedLoss() Code sees the union. A file or override never writes the tag; it writes the section under its kind as the key, `loss: {huber: {delta: 0.1}}` or `--loss.huber.delta 0.1`, or as a bare name, `loss: huber`, which is that kind with nothing set under it. The file and each override merge by the plain deep update of CORE-2; the loader records which layer last wrote each path and at each kinds field keeps the kind written last, dropping the others with a ConfigWarning that names both layers. Two kinds written in one place, and null at or under a kind, are ConfigErrors. A field with no kind written takes the kind of its default. The resolved and user dicts write the kind as key, so the resolved dict stays a fixed point. A before-validator on ConfigSection turns the key form into pydantic's tagged form and a wrap serializer turns it back; the tagged form is not a file format and load() refuses it. The schema check rejects at class definition what the contract cannot keep: a kinds field typed `| None` ("none of the kinds" is an empty variant, `kind: Literal["none"]`), a default that is None or a factory or not a variant, a non-variant arm, a kind named like the tag, and a union of sections inside a collection. Error paths name the kind as written, `loss.huber.delta`, and did-you-mean suggests kinds and their keys. --- .../src/mace_core/config/__init__.py | 9 +- .../mace-core/src/mace_core/config/base.py | 561 +++++++++-- .../tests/test_mace_core_config_kinds.py | 872 ++++++++++++++++++ 3 files changed, 1381 insertions(+), 61 deletions(-) create mode 100644 packages/mace-core/tests/test_mace_core_config_kinds.py diff --git a/packages/mace-core/src/mace_core/config/__init__.py b/packages/mace-core/src/mace_core/config/__init__.py index 9d6eaf542..9a507e7b1 100644 --- a/packages/mace-core/src/mace_core/config/__init__.py +++ b/packages/mace-core/src/mace_core/config/__init__.py @@ -7,8 +7,15 @@ from mace_core.config.base import ( ConfigError, ConfigSection, + ConfigWarning, ReforgeBaseConfig, read_config_file, ) -__all__ = ["ConfigError", "ConfigSection", "ReforgeBaseConfig", "read_config_file"] +__all__ = [ + "ConfigError", + "ConfigSection", + "ConfigWarning", + "ReforgeBaseConfig", + "read_config_file", +] diff --git a/packages/mace-core/src/mace_core/config/base.py b/packages/mace-core/src/mace_core/config/base.py index e9ebf81c4..8aa5b3d27 100644 --- a/packages/mace-core/src/mace_core/config/base.py +++ b/packages/mace-core/src/mace_core/config/base.py @@ -27,6 +27,43 @@ class TrainConfig(ReforgeBaseConfig): Nothing else feeds a config: no environment variables, no dotenv files, so a run is reproducible from its file and its command line alone. +A field may hold a section of one of several *kinds*: a discriminated union +whose tag field names the kind. + + class HuberLoss(ConfigSection): + kind: Literal["huber"] = "huber" + delta: float = 0.01 + + Loss = Annotated[WeightedLoss | HuberLoss, Field(discriminator="kind")] + + class TrainConfig(ReforgeBaseConfig): + loss: Loss = WeightedLoss() + +Code sees the union: `config.loss` is a `WeightedLoss` or a `HuberLoss`. A +file and the command line never write the tag; they write the kind as the +key the section sits under, or as a bare name, which is that kind with +nothing set under it: + + loss: huber --loss huber + loss: {huber: {delta: 0.1}} --loss.huber.delta 0.1 + +The file and each override merge in order. A bare name on top of a section +of the same kind keeps that section's keys; a different kind replaces it: +`--loss weighted` on top of a file with a huber section runs the weighted +loss and warns (`ConfigWarning`) that the huber section is ignored. Two +kinds written in one place, one file or one override, is an error, and so +is `null` at a kinds field or under a kind: a kinds field always holds a +kind. When none of the kinds is a valid choice, that is a kind too, an +empty variant such as `class NoLoss(ConfigSection): kind: Literal["none"]`, +written `loss: none`. A kinds field with no kind written takes the kind of +its default. The resolved and user dicts are written the same way, kind as +key. + +The tagged dict, `{kind: huber, delta: 0.1}`, is pydantic's internal form: +what `model_validate` takes and what `model_json_schema` describes. It is +not a file format; `load()` refuses it. Code builds sections as instances +(`HuberLoss(delta=0.1)`) and never meets either dict form. + Unknown keys are hard errors. A CLI key is checked against the schema's dotted paths before anything is built; a key in the file, or inside a JSON-valued override, is caught by pydantic (`extra="forbid"` on every level @@ -43,19 +80,35 @@ class TrainConfig(ReforgeBaseConfig): import difflib import json -from collections.abc import Iterator, Sequence +import warnings +from collections.abc import Callable, Iterator, Sequence +from itertools import cycle from pathlib import Path from types import UnionType -from typing import TYPE_CHECKING, Any, Union, get_args, get_origin +from typing import TYPE_CHECKING, Annotated, Any, Literal, Union, get_args, get_origin import tomli import yaml -from pydantic import BaseModel, ConfigDict, ValidationError +from pydantic import ( + BaseModel, + ConfigDict, + SerializerFunctionWrapHandler, + ValidationError, + model_serializer, + model_validator, +) +from pydantic.fields import FieldInfo if TYPE_CHECKING: from typing_extensions import Self -__all__ = ["ConfigError", "ConfigSection", "ReforgeBaseConfig", "read_config_file"] +__all__ = [ + "ConfigError", + "ConfigSection", + "ConfigWarning", + "ReforgeBaseConfig", + "read_config_file", +] #: Config file extensions this module reads, keyed to their parsers. _FILE_PARSERS = { @@ -65,17 +118,35 @@ class TrainConfig(ReforgeBaseConfig): ".json": json.loads, } +#: A dotted path as its parts; a list index is a part too. +_Path = tuple[str, ...] + +#: Where a value came from: the layer's index and its name ("the config file" +#: or the override as typed). Comparing two origins compares their order. +_Origin = tuple[int, str] + +#: What runs at a kinds field during a walk: gets the field's value (any +#: shape), the field and its path, returns the value to go on with. +_KindsAction = Callable[[Any, FieldInfo, _Path], Any] + + +# --------------------------------------------------------------------------- +# Schema introspection + def _sections_in( annotation: Any, inside: bool = False ) -> Iterator[tuple[type[BaseModel], bool]]: """Every section class a field annotation can hold (`Radial`, `Radial | None`, - `list[Radial]`), with whether it sits inside a dict/list/tuple, where an - error location has a key or index before the section's own field names.""" + `list[Radial]`, `Annotated[A | B, ...]`), with whether it sits inside a + dict/list/tuple, where an error location has a key or index before the + section's own field names.""" origin = get_origin(annotation) # `list` for `list[X]`; None for a plain class if origin is None: if isinstance(annotation, type) and issubclass(annotation, BaseModel): yield annotation, inside + elif origin is Annotated: + yield from _sections_in(get_args(annotation)[0], inside) elif origin in (Union, UnionType): for arg in get_args(annotation): yield from _sections_in(arg, inside) @@ -84,12 +155,75 @@ def _sections_in( yield from _sections_in(arg, True) +def _arms(annotation: Any) -> Iterator[Any]: + """The alternatives of a union, through `Annotated`; else the annotation.""" + origin = get_origin(annotation) + if origin is Annotated: + yield from _arms(get_args(annotation)[0]) + elif origin in (Union, UnionType): + for arg in get_args(annotation): + yield from _arms(arg) + else: + yield annotation + + def _section_of(annotation: Any) -> tuple[type[BaseModel] | None, bool]: """The one section a field can hold, and whether it is inside a collection.""" found = dict(_sections_in(annotation)) return next(iter(found.items())) if found else (None, False) +def _tag_of(field: FieldInfo) -> str | None: + """The tag field name of a kinds field, else None. Also found through an + outer union, `Annotated[...] | None`, which keeps it in the `Annotated` + metadata, so that `_check_schema` can reject that spelling by name.""" + found: Any = field.discriminator + if found is None and get_origin(field.annotation) in (Union, UnionType): + for arg in get_args(field.annotation): + if get_origin(arg) is Annotated: + for meta in get_args(arg)[1:]: + if isinstance(meta, FieldInfo) and meta.discriminator is not None: + found = meta.discriminator + return found if isinstance(found, str) else None + + +def _tag_values(section: type[BaseModel], tag: str) -> tuple[Any, ...]: + """The `Literal` values of a variant's tag field; empty if not a Literal.""" + tag_field = section.model_fields.get(tag) + if tag_field is None or get_origin(tag_field.annotation) is not Literal: + return () + return get_args(tag_field.annotation) + + +def _kinds_of(field: FieldInfo) -> dict[str, type[BaseModel]]: + """Kind name -> section class of a kinds field, in declaration order.""" + tag = _tag_of(field) + if tag is None: + return {} + return { + _tag_values(section, tag)[0]: section + for section, _ in _sections_in(field.annotation) + } + + +def _default_kind(field: FieldInfo) -> str | None: + """The kind of the field's default section; None for a required field.""" + tag = _tag_of(field) + default = field.get_default(call_default_factory=True) + if tag is None or not isinstance(default, BaseModel): + return None + return getattr(default, tag) + + +def _kinds_text(field: FieldInfo) -> str: + return ", ".join(_kinds_of(field)) + + +def _admits_none(annotation: Any) -> bool: + """`X | None`, `Optional[X]`, `Any`, `object`, also under `Annotated`.""" + return any(arm in (type(None), Any, object) for arm in _arms(annotation)) + + def _contains_a_set(annotation: Any) -> bool: """A bare `set`, a `set[X]`, or a set anywhere inside, e.g. `list[set[int]]`.""" if annotation in (set, frozenset) or get_origin(annotation) in (set, frozenset): @@ -97,16 +231,32 @@ def _contains_a_set(annotation: Any) -> bool: return any(_contains_a_set(arg) for arg in get_args(annotation)) +_NONE_IS_A_KIND = ( + "default to a variant, or for none of the kinds add an empty variant, " + 'kind: Literal["none"], and default to that' +) + + def _check_schema(model: type[BaseModel]) -> None: """Fail at class definition for a field shape the contract cannot keep. Each rule protects one guarantee: no sets (order is not stable across runs, so the export would not be a fixed point); no aliases or computed - fields (the export would not validate back); one section class per field - (a value must not become whichever alternative happens to accept it, and - the CLI needs one set of valid keys under a field); every section is a - `ConfigSection` (a plain `BaseModel` ignores unknown keys, so a typo - would vanish, and skips these checks). + fields (the export would not validate back); no `None` default on a type + that does not admit `None` (pydantic does not validate defaults, so the + export would not validate back; a default factory is not run here, so + what it returns is not checked); several sections under one field only + as kinds, i.e. a discriminated union, and not inside a collection (a + value must not become whichever alternative happens to accept it, and + the CLI addresses a collection only as a whole); nothing beside the + variants of a kinds field, not even `None` (a kinds field always holds a + kind; "none of them" is an empty variant, so it can be written, named + among the kinds and warned about like any other); a tag that is one + string other than the tag's own name (it is the key the kind is written + under); a default that is a variant written as an instance, not a + factory (its kind is the default kind, read without running anything); + every section is a `ConfigSection` (a plain `BaseModel` ignores unknown + keys, so a typo would vanish, and skips these checks). """ def reject(name: str, reason: str) -> None: @@ -122,12 +272,29 @@ def reject(name: str, reason: str) -> None: name, "is typed as a set; set order is not stable across runs. Use a list", ) - sections = {section for section, _ in _sections_in(field.annotation)} - if len(sections) > 1: + sections = dict(_sections_in(field.annotation)) + tag = _tag_of(field) + # `default` is undefined, not None, for a required field or a factory; + # a factory is not run here, so what it returns is not checked + if field.default is None and tag is not None: + reject(name, f"defaults to None, which is not a kind; {_NONE_IS_A_KIND}") + if field.default is None and not _admits_none(field.annotation): + reject( + name, "defaults to None but its type does not admit None; add | None" + ) + if len(sections) > 1 and any(sections.values()): + reject( + name, + "is a union of sections inside a dict, list or tuple; the CLI " + "addresses such a field only as a whole. Put the union in a " + "field of the section that is the element", + ) + if len(sections) > 1 and tag is None: reject( name, - "is a union of sections; give each alternative its own optional " - "field, e.g. `huber: HuberLoss | None`", + "is a union of sections without a discriminator; spell it " + 'Annotated[A | B, Field(discriminator="kind")] with a ' + '`kind: Literal["a"]` field in each', ) for section in sections: if not issubclass(section, ConfigSection): @@ -136,18 +303,62 @@ def reject(name: str, reason: str) -> None: f"holds {section.__name__}, which is not a ConfigSection; " f"subclass it", ) + if tag is None: + continue + others = [arm for arm in _arms(field.annotation) if arm not in sections] + if type(None) in others: + reject(name, f"admits None, which is not a kind; {_NONE_IS_A_KIND}") + if others: + reject( + name, + f"mixes its kinds with {getattr(others[0], '__name__', others[0])}; " + f"a kinds field holds its variants only", + ) + for section in sections: + values = _tag_values(section, tag) + if len(values) != 1 or not isinstance(values[0], str): + reject( + name, + f"has variant {section.__name__} whose {tag} must be a Literal " + f"of exactly one string; a config names the kind by it", + ) + if values[0] == tag: + reject( + name, + f"has variant {section.__name__} whose kind is named {tag!r} like " + f"the tag; a config could not tell the two apart. Rename it", + ) + example = f"{next(iter(sections)).__name__}()" + if field.default_factory is not None: + reject( + name, + f"has a default_factory; write the default as an instance, e.g. " + f"{example} (pydantic copies it per instance)", + ) + if not field.is_required() and not isinstance(field.default, tuple(sections)): + reject( + name, + f"has a default that is not one of its variants; write e.g. {example}", + ) class ConfigError(ValueError): """A config file or override the schema rejects. - Raised for a missing, unparsable or malformed file, an unknown key and an - override the CLI parser cannot make sense of. The message names the file - or the offending key by its dotted path, and suggests the nearest valid - key when there is a close match. + Raised for a missing, unparsable or malformed file, an unknown key or + kind, an override the CLI parser cannot make sense of, two kinds of one + section written in one place, and `null` at or under a kind. The + message names the file or the + offending key by its dotted path, and suggests the nearest valid key + when there is a close match. """ +class ConfigWarning(UserWarning): + """A section of one kind is ignored because a later layer selected another + kind. `warnings.simplefilter("error", ConfigWarning)` makes it an error.""" + + class ConfigSection(BaseModel): """A nested section of a configuration: a table in the file, a dotted prefix on the command line. Unknown keys are errors here too.""" @@ -159,6 +370,35 @@ def __pydantic_init_subclass__(cls, **kwargs: Any) -> None: super().__pydantic_init_subclass__(**kwargs) _check_schema(cls) + @model_validator(mode="before") + @classmethod + def _kinds_from_keys(cls, values: Any) -> Any: + """`{huber: {delta: 0.1}}` under a kinds field becomes the + `{kind: huber, delta: 0.1}` pydantic's discriminator reads.""" + if not isinstance(values, dict): + return values + values = dict(values) + for name, field in cls.model_fields.items(): + tag, value = _tag_of(field), values.get(name) + if tag is not None and isinstance(value, dict) and len(value) == 1: + ((kind, inner),) = value.items() + if isinstance(inner, dict) and tag not in value: + values[name] = {**inner, tag: kind} + return values + + @model_serializer(mode="wrap") + def _kinds_as_keys(self, handler: SerializerFunctionWrapHandler): + """The dump with every kinds field written kind-as-key. The kind is read + from the instance: an unset default tag is absent from a user dict. + No return annotation: pydantic would take it as the JSON schema.""" + dumped = handler(self) + for name, field in type(self).model_fields.items(): + tag = _tag_of(field) + if tag is not None and isinstance(dumped.get(name), dict): + inner = {key: v for key, v in dumped[name].items() if key != tag} + dumped[name] = {getattr(getattr(self, name), tag): inner} + return dumped + class ReforgeBaseConfig(ConfigSection): """Root of a configuration tree. Subclass it; nest `ConfigSection`s in it. @@ -183,18 +423,26 @@ def load( An override merges into the file, and into earlier overrides, like a section does: a dict-valued field gains or replaces entries, so an entry cannot be removed from the command line; a list-valued field - is replaced whole. + is replaced whole. Under a kinds field the kind written last wins + and the others are dropped with a `ConfigWarning`. - Raises `ConfigError` for an unknown key, an unreadable file or an - unparsable override, and pydantic's `ValidationError` for a value of - the wrong type. + Raises `ConfigError` for an unknown key or kind, an unreadable file + or an unparsable override, and pydantic's `ValidationError` for a + value of the wrong type. """ values: dict[str, Any] = {} if config_file is not None: values = read_config_file(config_file) - values = _apply_overrides(cls, values, cli_overrides) + layers = [(values, "the config file"), *_parse_overrides(cls, cli_overrides)] + merged: dict[str, Any] = {} + origins: dict[_Path, _Origin] = {} + for index, (layer, source) in enumerate(layers): + layer = _at_kinds(layer, cls, _name_as_mapping) + merged = _deep_update(merged, layer) + _record_origins(origins, layer, (index, source)) + merged = _at_kinds(merged, cls, _KindSelector(origins)) try: - return cls.model_validate(values) + return cls.model_validate(merged) except ValidationError as error: unknown = _unknown_key_messages(cls, error) if not unknown: @@ -203,8 +451,9 @@ def load( def to_resolved_dict(self) -> dict[str, Any]: """Every field, defaults included, as JSON-native values, in schema - order. Loading the result back and resolving again gives the same - dict. (TOML has no null, so a `None` can only go out as YAML or JSON.)""" + order, a kinds field as `{kind: {...}}`. Loading the result back and + resolving again gives the same dict. (TOML has no null, so a `None` + can only go out as YAML or JSON.)""" return self.model_dump(mode="json") def to_user_dict(self) -> dict[str, Any]: @@ -248,6 +497,10 @@ def read_config_file(path: str | Path) -> dict[str, Any]: return values +# --------------------------------------------------------------------------- +# Merging + + def _deep_update(base: dict[str, Any], update: dict[str, Any]) -> dict[str, Any]: """`base` overlaid with `update`, recursing where both hold a dict.""" merged = dict(base) @@ -259,14 +512,30 @@ def _deep_update(base: dict[str, Any], update: dict[str, Any]) -> dict[str, Any] return merged -def _apply_overrides( - model: type[BaseModel], values: dict[str, Any], cli_overrides: Sequence[str] -) -> dict[str, Any]: - """`values` with the `--a.b value` and `--a.b=value` pairs merged in, in order.""" +def _record_origins( + origins: dict[_Path, _Origin], values: Any, origin: _Origin, path: _Path = () +) -> None: + """Note `origin` for every path a layer writes, sections included.""" + items: Any = () + if isinstance(values, dict): + items = values.items() + elif isinstance(values, list): + items = enumerate(values) + for key, value in items: + origins[(*path, str(key))] = origin + _record_origins(origins, value, origin, (*path, str(key))) + + +def _parse_overrides( + model: type[BaseModel], cli_overrides: Sequence[str] +) -> list[tuple[dict[str, Any], str]]: + """Each `--a.b value` or `--a.b=value` pair as a mapping nested under its + path, with the override as typed, in order.""" valid = list(_dotted_paths(model)) valid_set = set(valid) unknown: list[str] = [] tokens = iter(cli_overrides) + layers = [] for token in tokens: option = token.removeprefix("--") @@ -278,10 +547,12 @@ def _apply_overrides( raise ConfigError( f"unknown config option {token!r}; overrides are written --key value" ) + source = token if value is None: value = next(tokens, None) if value is None: raise ConfigError(f"override --{name} is missing its value") + source = f"{token} {value}" if name not in valid_set: unknown.append(_unknown_key_message(name, valid)) @@ -295,20 +566,21 @@ def _apply_overrides( f"override --{name} is not valid JSON: {error}" ) from error - # Nested under its path, an override merges by the file's rule. override: Any = value for section in reversed(name.split(".")): override = {section: override} - values = _deep_update(values, override) + layers.append((override, source)) if unknown: raise ConfigError("\n".join(unknown)) - return values + return layers def _dotted_paths(model: type[BaseModel], prefix: str = "") -> Iterator[str]: - """Every field of the tree as a dotted path, sections included. + """Every field of the tree as a dotted path, sections included. A kinds + field lists each kind as a key with the kind's fields under it, without + the tag field: the key names the kind. A section inside a dict or list is not descended into: the CLI addresses such a field only as a whole, with a JSON value. @@ -316,11 +588,163 @@ def _dotted_paths(model: type[BaseModel], prefix: str = "") -> Iterator[str]: for name, field in model.model_fields.items(): path = f"{prefix}{name}" yield path + kinds = _kinds_of(field) + for kind, variant in kinds.items(): + yield f"{path}.{kind}" + for sub_path in _dotted_paths(variant, f"{path}.{kind}."): + if sub_path != f"{path}.{kind}.{_tag_of(field)}": + yield sub_path section, inside_collection = _section_of(field.annotation) - if section is not None and not inside_collection: + if section is not None and not inside_collection and not kinds: yield from _dotted_paths(section, f"{path}.") +# --------------------------------------------------------------------------- +# Kinds: a walk over the values against the schema, with an action at every +# kinds field. The action sees and returns the kind-as-key form. + + +def _at_kinds( + values: Any, section: type[BaseModel] | None, action: _KindsAction, path: _Path = () +) -> Any: + """`values`, a mapping for `section`, with `action` applied at each kinds + field and every section under it walked in turn.""" + if section is None or not isinstance(values, dict): + return values + out = dict(values) + for key, value in values.items(): + field = section.model_fields.get(key) + if field is None: + continue + kinds = _kinds_of(field) + if not kinds: + out[key] = _under(value, field.annotation, action, (*path, key)) + continue + value = action(value, field, (*path, key)) + if isinstance(value, dict): + value = { + kind: _at_kinds(inner, kinds.get(kind), action, (*path, key, kind)) + for kind, inner in value.items() + } + out[key] = value + return out + + +def _under(value: Any, annotation: Any, action: _KindsAction, path: _Path) -> Any: + """`value` with `_at_kinds` applied to every section the annotation reaches + through unions, dicts, lists and tuples. A value whose shape the annotation + does not describe is returned as is, for pydantic to report.""" + origin = get_origin(annotation) + if origin is None: + return _at_kinds(value, _section_of(annotation)[0], action, path) + if origin is Annotated: + return _under(value, get_args(annotation)[0], action, path) + if origin in (Union, UnionType): # at most one arm takes a dict or a list + for arm in get_args(annotation): + value = _under(value, arm, action, path) + return value + if origin is dict and isinstance(value, dict): + value_type = get_args(annotation)[1] + return { + k: _under(v, value_type, action, (*path, str(k))) for k, v in value.items() + } + if origin in (list, tuple) and isinstance(value, list): + item_types = [a for a in get_args(annotation) if a is not Ellipsis] + return [ + _under(v, t, action, (*path, str(i))) + for i, (v, t) in enumerate(zip(value, cycle(item_types), strict=False)) + ] + return value + + +def _shown(value: Any) -> str: + """A value as the user could have written it; a date or a YAML set as text.""" + return json.dumps(value, default=str) + + +def _name_as_mapping(value: Any, field: FieldInfo, path: _Path) -> Any: + """A bare kind name is that kind with its defaults, `{huber: {}}`, so that + it merges into an earlier section of the same kind instead of replacing it.""" + return {value: {}} if isinstance(value, str) else value + + +class _KindSelector: + """At a kinds field after the merge: keep the kind written last, drop the + others with a warning, fill in the default kind, and check the shape.""" + + def __init__(self, origins: dict[_Path, _Origin]) -> None: + self.origins = origins + + def __call__(self, value: Any, field: FieldInfo, path: _Path) -> Any: + dotted = ".".join(path) + tag, kinds, named = _tag_of(field), _kinds_of(field), _kinds_text(field) + if value is None: + raise ConfigError( + f"{dotted} does not take null; write a kind, one of {named}" + ) + if not isinstance(value, dict): + raise ConfigError( + f"{dotted} must be the name of a kind or a mapping under one, " + f"one of {named}; got {_shown(value)}" + ) + if tag in value: + raise ConfigError( + f"{dotted}.{tag} is not a key; write the kind as the key the " + f"section sits under, {dotted}: {{{_shown(value[tag])}: {{...}}}}" + ) + nulled = [k for k, inner in value.items() if inner is None and k in kinds] + if nulled: + raise ConfigError( + f"{dotted}.{nulled[0]} does not take null; set the keys wanted " + f"under it, or write another kind" + ) + if not value: + default = _default_kind(field) + if default is None: + raise ConfigError(f"{dotted} needs a kind; one of {named}") + return {default: {}} + present = value + origin_of = {k: self.origins[(*path, str(k))] for k in present} + by_origin = sorted(present, key=origin_of.__getitem__) + kind = by_origin[-1] + index, source = origin_of[kind] + tied = [k for k in by_origin if origin_of[k][0] == index] + if len(tied) > 1: + raise ConfigError( + f"{dotted} is given as several kinds ({', '.join(tied)}) in " + f"{source}; keep one" + ) + if kind not in kinds: + valid = [f"{dotted}.{k}" for k in kinds] + raise ConfigError( + _unknown_key_message(f"{dotted}.{kind}", valid) + + f"; the kinds of {dotted} are {named}" + ) + for loser in by_origin[:-1]: + warnings.warn( + f"{dotted}.{loser} from {origin_of[loser][1]} is " + f"ignored: {source} selects {dotted}.{kind}", + ConfigWarning, + stacklevel=2, + ) + inner = present[kind] + if not isinstance(inner, dict): + raise ConfigError( + f"{dotted}.{kind} must be a mapping of the kind's keys; " + f"got {_shown(inner)}" + ) + if tag in inner: + raise ConfigError( + f"{dotted}.{kind}.{tag} is not a key; the kind is given by the " + f"key {kind!r}" + ) + return {kind: inner} + + +# --------------------------------------------------------------------------- +# Error messages + + def _unknown_key_message(key: str, candidates: Sequence[str]) -> str: message = f"unknown config key {key!r}" closest = difflib.get_close_matches(key, candidates, n=1) @@ -329,38 +753,55 @@ def _unknown_key_message(key: str, candidates: Sequence[str]) -> str: return message +def _locate( + model: type[BaseModel], location: Sequence[Any] +) -> tuple[list[str], list[str]]: + """Pydantic's error location as the names of a dotted path, and the keys + valid where it ends. A field name moves into its section; a kind moves + into its variant, whose tag field is not offered since the key names the + kind; a dict key or list index stays in the section; the class name + pydantic inserts under a `Section | scalar` field is dropped.""" + section: type[BaseModel] | None = model + kinds: dict[str, type[BaseModel]] | None = None + hidden: str | None = None + inside_collection = False + names = [] + for part in map(str, location): + if kinds is not None: + names.append(part) + section, kinds = kinds.get(part), None + continue + if inside_collection: + names.append(part) + inside_collection = False + continue + if section is not None and part == section.__name__: + continue + names.append(part) + if section is not None and part in section.model_fields: + field = section.model_fields[part] + hidden = _tag_of(field) + if hidden is not None: + kinds = _kinds_of(field) + else: + section, inside_collection = _section_of(field.annotation) + else: + section = None + candidates = list(section.model_fields) if section is not None else [] + return names, [c for c in candidates if c != hidden] + + def _unknown_key_messages(model: type[BaseModel], error: ValidationError) -> list[str]: """Pydantic's unknown-key errors as messages that name the full dotted - path and the closest valid key at that level. - - The error location is walked against the schema to find the section whose - fields are the candidates: a field name moves into its section, a dict key - or list index stays in it, and the tag pydantic inserts for a - `Section | scalar` field (the class name) is dropped from the path. - """ + path and the closest valid key at that level.""" messages = [] for item in error.errors(): # pydantic's error code for a key that matches no field (extra="forbid"). # Every other code, e.g. a wrong type, is left for `load()` to re-raise. if item["type"] != "extra_forbidden": continue - *location, key = (str(part) for part in item["loc"]) - section: type[BaseModel] | None = model - inside_collection = False - names = [] - for part in location: - if section is not None and part == section.__name__: - continue - names.append(part) - if inside_collection: - inside_collection = False - elif section is not None and part in section.model_fields: - section, inside_collection = _section_of( - section.model_fields[part].annotation - ) - else: - section = None - candidates = list(section.model_fields) if section is not None else [] + *location, key = item["loc"] + names, candidates = _locate(model, location) prefix = "".join(f"{name}." for name in names) messages.append( _unknown_key_message(f"{prefix}{key}", [f"{prefix}{c}" for c in candidates]) diff --git a/packages/mace-core/tests/test_mace_core_config_kinds.py b/packages/mace-core/tests/test_mace_core_config_kinds.py new file mode 100644 index 000000000..2e3d7a24e --- /dev/null +++ b/packages/mace-core/tests/test_mace_core_config_kinds.py @@ -0,0 +1,872 @@ +"""A field of several kinds of section (a discriminated union) under the file +and dotted-override contract. A config writes the kind as the key the section +sits under (`loss: {huber: {delta: 0.1}}`, `--loss.huber.delta 0.1`) or as a +bare name for the kind with its defaults; code sees the union. The file and +every override merge in order, the kind written last wins, the others are +dropped with a warning, two kinds in one place is an error.""" + +import json +import re +import warnings +from typing import Annotated, Any, Literal + +import pytest +from mace_core.config import ( + ConfigError, + ConfigSection, + ConfigWarning, + ReforgeBaseConfig, +) +from pydantic import BaseModel, Field, ValidationError + +# --------------------------------------------------------------------------- +# The schema: a loss of three kinds, one of which holds a field of two kinds; +# the same choice with "none of them" as a fourth kind and the default, and +# inside the values of a dict and the items of a list. + + +class SubX(ConfigSection): + kind: Literal["x"] = "x" + a: float = 1.0 + + +class SubY(ConfigSection): + kind: Literal["y"] = "y" + b: float = 2.0 + + +class Weighted(ConfigSection): + kind: Literal["weighted"] = "weighted" + stress_weight: float = 0.0 + + +class Huber(ConfigSection): + kind: Literal["huber"] = "huber" + delta: float = 0.01 + sub: Annotated[SubX | SubY, Field(discriminator="kind")] = SubX() + + +class Universal(ConfigSection): + kind: Literal["universal"] = "universal" + huber_delta: float = 0.01 + + +class Plain(ConfigSection): + p: int = 0 + + +Choice = Annotated[Weighted | Huber | Universal, Field(discriminator="kind")] + + +class NoChoice(ConfigSection): + """None of the other kinds is a kind too.""" + + kind: Literal["none"] = "none" + + +OptChoice = Annotated[ + Weighted | Huber | Universal | NoChoice, Field(discriminator="kind") +] + + +class HeadSection(ConfigSection): + loss: Choice = Weighted() + + +class LossConfig(ReforgeBaseConfig): + energy_weight: float = 1.0 + choice: Choice = Weighted() + opt: OptChoice = NoChoice() + heads: dict[str, Plain] = Field(default_factory=dict) + per_head: dict[str, HeadSection] = Field(default_factory=dict) + layers: list[HeadSection] = Field(default_factory=list) + + +class HuberRequired(ConfigSection): + kind: Literal["huber"] = "huber" + path: str + delta: float = 0.01 + + +class RequiredConfig(ReforgeBaseConfig): + choice: Annotated[Weighted | HuberRequired, Field(discriminator="kind")] = ( + Weighted() + ) + + +# --------------------------------------------------------------------------- +# Expectations. A row is (file values, command line, expectation); the +# expectation is a predicate on the loaded config, optionally with the +# warnings the load must emit, or an error spec. + +KINDS = "weighted, huber, universal" +OPT_KINDS = "weighted, huber, universal, none" + + +class Raises: + def __init__(self, error_type, *fragments): + self.error_type = error_type + self.fragments = fragments + + +class Warns: + """A predicate plus the exact warning messages, in order.""" + + def __init__(self, predicate, *messages): + self.predicate = predicate + self.messages = messages + + +def error(*fragments): + """A `ConfigError` whose message holds every fragment.""" + return Raises(ConfigError, *fragments) + + +def ignored(loser, source, winner_source, winner): + return f"{loser} from {source} is ignored: {winner_source} selects {winner}" + + +def huber(delta=0.01, sub=SubX, **sub_fields): + return lambda c: ( + isinstance(c.choice, Huber) + and c.choice.delta == delta + and isinstance(c.choice.sub, sub) + and all(getattr(c.choice.sub, k) == v for k, v in sub_fields.items()) + ) + + +def weighted(stress_weight=0.0): + return lambda c: ( + isinstance(c.choice, Weighted) and c.choice.stress_weight == stress_weight + ) + + +def universal(huber_delta=0.01): + return lambda c: ( + isinstance(c.choice, Universal) and c.choice.huber_delta == huber_delta + ) + + +def opt(kind, **fields): + return lambda c: ( + isinstance(c.opt, kind) + and all(getattr(c.opt, k) == v for k, v in fields.items()) + ) + + +def per_head(name, kind, **fields): + return lambda c: ( + isinstance(c.per_head[name].loss, kind) + and all(getattr(c.per_head[name].loss, k) == v for k, v in fields.items()) + ) + + +def layer(index, kind, **fields): + return lambda c: ( + isinstance(c.layers[index].loss, kind) + and all(getattr(c.layers[index].loss, k) == v for k, v in fields.items()) + ) + + +HUBER_FILE = {"choice": {"huber": {"delta": 0.5}}} + +SWITCHING = { + "file kind, no cli": (HUBER_FILE, "", huber(0.5)), + "cli name switches kind, file section dropped with a warning": ( + HUBER_FILE, + "--choice weighted", + Warns( + weighted(), + ignored( + "choice.huber", + "the config file", + "--choice weighted", + "choice.weighted", + ), + ), + ), + "cli dotted key switches kind": ( + HUBER_FILE, + "--choice.weighted.stress_weight 5", + Warns( + weighted(5.0), + ignored( + "choice.huber", + "the config file", + "--choice.weighted.stress_weight 5", + "choice.weighted", + ), + ), + ), + "cli json switches kind": ( + HUBER_FILE, + '--choice {"universal": {"huber_delta": 3}}', + Warns( + universal(3.0), + ignored( + "choice.huber", + "the config file", + '--choice {"universal": {"huber_delta": 3}}', + "choice.universal", + ), + ), + ), + "cli name of the file's kind keeps the file's keys": ( + HUBER_FILE, + "--choice huber", + huber(0.5), + ), + "cli key of the file's kind merges": ( + HUBER_FILE, + "--choice.huber.sub y", + huber(0.5, SubY), + ), + "switch away and back keeps the file's keys": ( + HUBER_FILE, + "--choice weighted --choice.huber.sub y", + Warns( + huber(0.5, SubY), + ignored( + "choice.weighted", + "--choice weighted", + "--choice.huber.sub y", + "choice.huber", + ), + ), + ), + "two kinds on the cli, the later wins": ( + {}, + "--choice.huber.delta 2 --choice.weighted.stress_weight 5", + Warns( + weighted(5.0), + ignored( + "choice.huber", + "--choice.huber.delta 2", + "--choice.weighted.stress_weight 5", + "choice.weighted", + ), + ), + ), + "no file, default kind": ({}, "", weighted()), + "no file, cli key of another kind": ({}, "--choice.huber.delta 2", huber(2.0)), + "empty section is the default kind": ({"choice": {}}, "", weighted()), + "empty section with cli key": ( + {"choice": {}}, + "--choice.huber.delta 2", + huber(2.0), + ), + "bare name in the file": ({"choice": "huber"}, "", huber()), + "bare name in the file, cli key of that kind": ( + {"choice": "huber"}, + "--choice.huber.delta 2", + huber(2.0), + ), + "bare name in the file, cli other kind": ( + {"choice": "huber"}, + "--choice.weighted.stress_weight 5", + Warns( + weighted(5.0), + ignored( + "choice.huber", + "the config file", + "--choice.weighted.stress_weight 5", + "choice.weighted", + ), + ), + ), + "= form": (HUBER_FILE, "--choice.huber.delta=2", huber(2.0)), + "three kinds in a row": ( + HUBER_FILE, + "--choice weighted --choice universal", + Warns( + universal(), + ignored( + "choice.huber", + "the config file", + "--choice universal", + "choice.universal", + ), + ignored( + "choice.weighted", + "--choice weighted", + "--choice universal", + "choice.universal", + ), + ), + ), +} + +ERRORS = { + "two kinds in the file": ( + {"choice": {"huber": {}, "weighted": {}}}, + "", + error( + "choice is given as several kinds (huber, weighted) in the config " + "file; keep one" + ), + ), + "two kinds in one json override": ( + {}, + '--choice {"huber": {}, "weighted": {}}', + error("choice is given as several kinds (huber, weighted) in", "keep one"), + ), + "unknown kind in the file": ( + {"choice": {"hubr": {}}}, + "", + error( + "unknown config key 'choice.hubr'; did you mean 'choice.huber'?; " + f"the kinds of choice are {KINDS}" + ), + ), + "unknown bare name in the file": ( + {"choice": "hubr"}, + "", + error("unknown config key 'choice.hubr'; did you mean 'choice.huber'?"), + ), + "unknown kind on the cli": ( + {}, + "--choice hubr", + error("unknown config key 'choice.hubr'; did you mean 'choice.huber'?"), + ), + "unknown dotted kind on the cli": ( + {}, + "--choice.hubr.delta 2", + error( + "unknown config key 'choice.hubr.delta'; did you mean 'choice.huber.delta'?" + ), + ), + "key of another kind on the cli is unknown": ( + {}, + "--choice.delta 2", + error("unknown config key 'choice.delta'; did you mean 'choice.huber.delta'?"), + ), + "key of another kind in the file": ( + {"choice": {"weighted": {"delta": 2}}}, + "", + error( + "unknown config key 'choice.weighted.delta'; did you mean " + "'choice.weighted.stress_weight'?" + ), + ), + "the tag is not a key, on the cli": ( + {}, + "--choice.huber.kind huber", + error("unknown config key 'choice.huber.kind'"), + ), + "the tag is not a key, in the file": ( + {"choice": {"huber": {"kind": "weighted"}}}, + "", + error("choice.huber.kind is not a key; the kind is given by the key 'huber'"), + ), + "the tagged form is refused with the key form as the fix": ( + {"choice": {"kind": "huber", "delta": 2}}, + "", + error( + "choice.kind is not a key; write the kind as the key the section " + 'sits under, choice: {"huber": {...}}' + ), + ), + "the tagged form on the cli": ( + {}, + '--choice {"kind": "huber"}', + error("choice.kind is not a key"), + ), + "a list at a kinds field": ( + {"choice": [1]}, + "", + error( + "choice must be the name of a kind or a mapping under one, one of " + f"{KINDS}; got [1]" + ), + ), + "a number at a kinds field": ( + {"choice": 3}, + "", + error( + "choice must be the name of a kind or a mapping under one, one of " + f"{KINDS}; got 3" + ), + ), + "a scalar under a kind": ( + {"choice": {"huber": 3}}, + "", + error("choice.huber must be a mapping of the kind's keys; got 3"), + ), + "null at a kinds field that does not admit it": ( + {"choice": None}, + "", + error(f"choice does not take null; write a kind, one of {KINDS}"), + ), + "a dropped section is not validated": ( + {"choice": {"huber": {"nonsense": 1}}}, + "--choice weighted", + Warns( + weighted(), + ignored( + "choice.huber", + "the config file", + "--choice weighted", + "choice.weighted", + ), + ), + ), +} + +NULL = { + "null override then a key starts the section afresh": ( + HUBER_FILE, + "--choice null --choice.huber.sub y", + huber(0.01, SubY), + ), + "null under a kind, on the cli": ( + HUBER_FILE, + "--choice.huber null", + error( + "choice.huber does not take null; set the keys wanted under it, " + "or write another kind" + ), + ), + "null under an unknown kind is an unknown kind": ( + {"choice": {"hubr": None}}, + "", + error("unknown config key 'choice.hubr'; did you mean 'choice.huber'?"), + ), + "null under a kind, in the file": ( + {"choice": {"huber": {"delta": 0.5}, "weighted": None}}, + "", + error("choice.weighted does not take null"), + ), +} + +NONE_KIND = { + "none kind, absent": ({}, "", opt(NoChoice)), + "none kind, by name": ({}, "--opt none", opt(NoChoice)), + "none kind, cli key": ({}, "--opt.huber.delta 2", opt(Huber, delta=2.0)), + "none kind, file null": ( + {"opt": None}, + "", + error(f"opt does not take null; write a kind, one of {OPT_KINDS}"), + ), + "none kind, cli null after the file": ( + {"opt": {"huber": {}}}, + "--opt null", + error(f"opt does not take null; write a kind, one of {OPT_KINDS}"), + ), + "none kind, back to none with a warning": ( + {"opt": {"huber": {}}}, + "--opt none", + Warns( + opt(NoChoice), + ignored("opt.huber", "the config file", "--opt none", "opt.none"), + ), + ), +} + +NESTED = { + "nested kind by name": (HUBER_FILE, "--choice.huber.sub y", huber(0.5, SubY)), + "nested kind by key": ( + HUBER_FILE, + "--choice.huber.sub.y.b 7", + huber(0.5, SubY, b=7.0), + ), + "nested kind in the file, cli key of it": ( + {"choice": {"huber": {"sub": {"y": {"b": 3}}}}}, + "--choice.huber.sub.y.b 7", + huber(0.01, SubY, b=7.0), + ), + "nested switch warns": ( + {"choice": {"huber": {"sub": {"y": {"b": 3}}}}}, + "--choice.huber.sub x", + Warns( + huber(0.01, SubX), + ignored( + "choice.huber.sub.y", + "the config file", + "--choice.huber.sub x", + "choice.huber.sub.x", + ), + ), + ), + "nested empty section is its default kind": ( + {"choice": {"huber": {"sub": {}}}}, + "", + huber(0.01, SubX), + ), + "outer switch drops the nested section silently": ( + {"choice": {"huber": {"sub": {"y": {"b": 3}}}}}, + "--choice weighted", + Warns( + weighted(), + ignored( + "choice.huber", + "the config file", + "--choice weighted", + "choice.weighted", + ), + ), + ), + "unknown nested kind": ( + {"choice": {"huber": {"sub": {"z": {}}}}}, + "", + error( + "unknown config key 'choice.huber.sub.z'", + "the kinds of choice.huber.sub are x, y", + ), + ), + "unknown key under a nested kind": ( + {"choice": {"huber": {"sub": {"y": {"a": 1}}}}}, + "", + error( + "unknown config key 'choice.huber.sub.y.a'; did you mean " + "'choice.huber.sub.y.b'?" + ), + ), +} + +COLLECTIONS = { + "kind inside a dict value, by name": ( + {"per_head": {"h": {"loss": "huber"}}}, + "", + per_head("h", Huber), + ), + "kind inside a dict value, by key": ( + {"per_head": {"h": {"loss": {"huber": {"delta": 2}}}}}, + "", + per_head("h", Huber, delta=2.0), + ), + "kind inside a dict value, json merges the entry": ( + {"per_head": {"h": {"loss": {"huber": {"delta": 2}}}}}, + '--per_head {"h": {"loss": "huber"}}', + per_head("h", Huber, delta=2.0), + ), + "kind inside a dict value, json switches with a warning": ( + {"per_head": {"h": {"loss": {"huber": {"delta": 2}}}}}, + '--per_head {"h": {"loss": "weighted"}}', + Warns( + per_head("h", Weighted), + ignored( + "per_head.h.loss.huber", + "the config file", + '--per_head {"h": {"loss": "weighted"}}', + "per_head.h.loss.weighted", + ), + ), + ), + "kind inside a dict value, empty is the default": ( + {"per_head": {"h": {"loss": {}}}}, + "", + per_head("h", Weighted), + ), + "unknown key inside a dict value names the kind": ( + {"per_head": {"h": {"loss": {"huber": {"stress_weight": 1}}}}}, + "", + error( + "unknown config key 'per_head.h.loss.huber.stress_weight'; did you mean " + "'per_head.h.loss.huber.delta'?" + ), + ), + "kind inside a list item": ( + {"layers": [{"loss": "huber"}, {"loss": {"universal": {"huber_delta": 3}}}]}, + "", + lambda c: layer(0, Huber)(c) and layer(1, Universal, huber_delta=3.0)(c), + ), + "list replaced whole by json": ( + {"layers": [{"loss": "huber"}]}, + '--layers [{"loss": {"weighted": {"stress_weight": 5}}}]', + lambda c: len(c.layers) == 1 and layer(0, Weighted, stress_weight=5.0)(c), + ), + "unknown key inside a list item": ( + {"layers": [{"loss": {"huber": {"stress_weight": 1}}}]}, + "", + error( + "unknown config key 'layers.0.loss.huber.stress_weight'; did you mean " + "'layers.0.loss.huber.delta'?" + ), + ), + "list item cannot be addressed by a dotted key": ( + {}, + "--layers.0.loss huber", + error("unknown config key 'layers.0.loss'"), + ), +} + +OTHER = { + "keys of other fields are untouched by a switch": ( + {"energy_weight": 3.0, **HUBER_FILE}, + "--choice weighted", + Warns( + lambda c: c.energy_weight == 3.0 and weighted()(c), + ignored( + "choice.huber", + "the config file", + "--choice weighted", + "choice.weighted", + ), + ), + ), + "a required key of the kind supplied by the cli": ( + {"choice": {"huber": {}}}, + "--choice.huber.path p", + lambda c: isinstance(c.choice, HuberRequired) and c.choice.path == "p", + ), + "a required key of the kind missing": ( + {"choice": "huber"}, + "", + Raises(ValidationError, "choice.huber.path", "Field required"), + ), +} + +ROWS = {**SWITCHING, **ERRORS, **NULL, **NONE_KIND, **NESTED, **COLLECTIONS} + + +def split_cli(cli): + """`--a.b 1 --c {"d": 2}` as argv: a JSON value keeps its spaces.""" + argv = [] + for chunk in filter(None, re.split(r"\s+(?=--)", cli)): + argv.extend(chunk.split(" ", 1)) + return argv + + +def load(tmp_path, root, file_values, cli): + path = tmp_path / "config.json" + path.write_text(json.dumps(file_values)) + return root.load(path, cli_overrides=split_cli(cli)) + + +def check(tmp_path, root, file_values, cli, expected): + if isinstance(expected, Raises): + with pytest.raises(expected.error_type) as info, warnings.catch_warnings(): + warnings.simplefilter("ignore", ConfigWarning) + load(tmp_path, root, file_values, cli) + for fragment in expected.fragments: + assert fragment in str(info.value), str(info.value) + return + predicate, messages = expected, () + if isinstance(expected, Warns): + predicate, messages = expected.predicate, expected.messages + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + config = load(tmp_path, root, file_values, cli) + assert [str(w.message) for w in caught] == list(messages) + assert all(issubclass(w.category, ConfigWarning) for w in caught) + assert predicate(config), config + + +@pytest.mark.parametrize("row", ROWS, ids=ROWS) +def test_kind_selection(tmp_path, row): + check(tmp_path, LossConfig, *ROWS[row]) + + +@pytest.mark.parametrize("row", OTHER, ids=OTHER) +def test_kind_selection_on_other_roots(tmp_path, row): + file_values, cli, expected = OTHER[row] + root = RequiredConfig if "required" in row else LossConfig + check(tmp_path, root, file_values, cli, expected) + + +@pytest.mark.parametrize( + ("extension", "text", "shown"), + [ + (".toml", "choice = 2020-01-01\n", '"2020-01-01"'), + (".yaml", "choice:\n huber: 2020-01-01\n", '"2020-01-01"'), + (".yaml", "choice: !!set {huber: null}\n", "\"{'huber'}\""), + ], +) +def test_a_value_json_cannot_show_is_still_a_config_error( + tmp_path, extension, text, shown +): + path = tmp_path / f"config{extension}" + path.write_text(text) + with pytest.raises(ConfigError, match=rf"got {shown}$"): + LossConfig.load(path) + + +@pytest.mark.parametrize("extension", [".toml", ".yaml"]) +def test_kinds_load_from_every_file_format(tmp_path, extension): + text = { + ".toml": "[choice.huber]\ndelta = 0.5\n[choice.huber.sub.y]\nb = 3\n" + "[opt]\nuniversal = {}\n", + ".yaml": "choice:\n huber:\n delta: 0.5\n sub:\n y:\n" + " b: 3\nopt: universal\n", + }[extension] + path = tmp_path / f"config{extension}" + path.write_text(text) + config = LossConfig.load(path) + assert huber(0.5, SubY, b=3.0)(config) and isinstance(config.opt, Universal) + + +# --------------------------------------------------------------------------- +# The dicts: kind as key, a fixed point, and only what was set. + + +@pytest.mark.parametrize( + "cli", + [ + "", + "--choice huber", + "--choice.huber.sub.y.b 7 --opt universal", + '--per_head {"h": {"loss": "huber"}} --layers [{"loss": {"universal": {}}}]', + ], +) +def test_resolved_dict_is_a_fixed_point_and_writes_the_kind_as_key(tmp_path, cli): + config = load(tmp_path, LossConfig, {}, cli) + resolved = config.to_resolved_dict() + kind = type(config.choice).model_fields["kind"].default + assert list(resolved["choice"]) == [kind] + assert "kind" not in resolved["choice"][kind] + reloaded = load(tmp_path, LossConfig, resolved, "") + assert reloaded.to_resolved_dict() == resolved + assert reloaded == config + + +def test_resolved_dict_of_the_defaults(): + assert LossConfig().to_resolved_dict() == { + "energy_weight": 1.0, + "choice": {"weighted": {"stress_weight": 0.0}}, + "opt": {"none": {}}, + "heads": {}, + "per_head": {}, + "layers": [], + } + + +def test_user_dict_holds_the_kind_and_only_what_was_set(tmp_path): + config = load(tmp_path, LossConfig, {"choice": "huber"}, "--choice.huber.delta 2") + assert config.to_user_dict() == {"choice": {"huber": {"delta": 2.0}}} + # A kind chosen by default is what ran, so it is set. + assert load(tmp_path, LossConfig, {"choice": {}}, "").to_user_dict() == { + "choice": {"weighted": {}} + } + # Built in code, the tag is an unset default; the kind is still the key. + assert LossConfig(choice=Huber(delta=2)).to_user_dict() == { + "choice": {"huber": {"delta": 2.0}} + } + + +def test_json_schema_is_produced_in_both_modes(): + for mode in ("validation", "serialization"): + schema = LossConfig.model_json_schema(mode=mode) + assert set(schema["properties"]) == set(LossConfig.model_fields), mode + + +def test_code_sees_the_union_and_may_construct_it_either_way(): + by_key = {"choice": {"huber": {"delta": 2}}} + by_tag = {"choice": {"kind": "huber", "delta": 2}} + assert LossConfig(choice=Huber(delta=2)).choice == Huber(delta=2) + assert LossConfig.model_validate(by_key).choice == Huber(delta=2) + assert LossConfig.model_validate(by_tag).choice == Huber(delta=2) + + +# --------------------------------------------------------------------------- +# Schema rules. + + +class NotASection(BaseModel): + kind: Literal["plain"] = "plain" + + +def test_union_shapes_the_contract_cannot_keep_are_rejected_at_class_definition(): + class IntTag(ConfigSection): + kind: Literal[1] = 1 + + class TwoTags(ConfigSection): + kind: Literal["a", "b"] = "a" + + class NamedLikeTheTag(ConfigSection): + kind: Literal["kind"] = "kind" + + shapes = { + r"choice defaults to None, which is not a kind; default to a variant": ( + Choice, + None, + ), + r"choice admits None, which is not a kind; default to a variant, or for none": ( + Choice | None, + Weighted(), + ), + r"choice has variant NamedLikeTheTag whose kind is named 'kind' like the tag": ( + Annotated[Weighted | NamedLikeTheTag, Field(discriminator="kind")], + Weighted(), + ), + r"choice is a union of sections without a discriminator": ( + Weighted | Huber, + Weighted(), + ), + r"choice is a union of sections inside a dict, list or tuple": ( + list[Choice], + [], + ), + r"choice holds NotASection, which is not a ConfigSection": ( + Annotated[Weighted | NotASection, Field(discriminator="kind")], + Weighted(), + ), + r"choice has variant IntTag whose kind must be a Literal of exactly one": ( + Annotated[Weighted | IntTag, Field(discriminator="kind")], + Weighted(), + ), + r"choice has variant TwoTags whose kind must be a Literal of exactly one": ( + Annotated[Weighted | TwoTags, Field(discriminator="kind")], + Weighted(), + ), + r"choice has a default that is not one of its variants; write e.g. Weighted": ( + Choice, + Plain(), + ), + r"choice has a default_factory; write the default as an instance": ( + Weighted | Huber, + Field(discriminator="kind", default_factory=Weighted), + ), + r"choice mixes its kinds with dict; a kinds field holds its variants only": ( + Choice | dict[str, int], + Weighted(), + ), + } + for message, (annotation, default) in shapes.items(): + with pytest.raises(TypeError, match=message): + type( + "Bad", + (ConfigSection,), + {"__annotations__": {"choice": annotation}, "choice": default}, + ) + + +def test_a_required_kinds_field_is_accepted(): + factory_calls = [] + + class Required(ReforgeBaseConfig): + choice: Choice + anything: Any = None + made: int = Field(default_factory=lambda: factory_calls.append(1) or 1) + + assert factory_calls == [], "the schema check must not run default factories" + with pytest.raises(ConfigError, match="choice needs a kind; one of"): + Required.load(cli_overrides=["--choice", "{}"]) + assert isinstance(Required.load(cli_overrides=["--choice", "huber"]).choice, Huber) + + +def test_a_kind_named_like_a_class_still_reads_as_a_kind_in_error_paths(): + class Root(ConfigSection): + kind: Literal["Root"] = "Root" + q: int = 0 + + class Config(ReforgeBaseConfig): + choice: Annotated[Weighted | Root, Field(discriminator="kind")] = Weighted() + + with pytest.raises( + ConfigError, match=r"'choice\.Root\.qq'; did you mean 'choice\.Root\.q'" + ): + Config.load(cli_overrides=["--choice", '{"Root": {"qq": 1}}']) + + +def test_a_collection_key_named_like_the_element_class_stays_in_error_paths(): + with pytest.raises( + ConfigError, match=r"'heads\.Plain\.q'; did you mean 'heads\.Plain\.p'" + ): + LossConfig.load(cli_overrides=["--heads", '{"Plain": {"q": 1}}']) + + +def test_warning_can_be_turned_into_an_error(tmp_path): + with warnings.catch_warnings(): + warnings.simplefilter("error", ConfigWarning) + with pytest.raises(ConfigWarning, match=r"choice\.huber from the config file"): + load(tmp_path, LossConfig, HUBER_FILE, "--choice weighted") From b312a933738f921ce8a6bc5514baf64ef7172cf4 Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Mon, 21 Sep 2026 14:01:11 +0200 Subject: [PATCH 04/14] Check a config schema behind a forward reference when load resolves it (CORE-2 follow-up, #1556) The schema check ran once, at class definition. A field whose annotation still named a class defined later was an unresolved forward reference then, opaque to every rule: `q: "Later | None" = None` was rejected with an error that contradicted itself, and a plain BaseModel behind such a name passed the ConfigSection rule and silently dropped unknown keys. Pydantic's rebuild of an outer model does not rebuild an inner one, so a rebuild hook alone would miss a nested section. An incomplete class is now skipped at definition; `load` completes and checks every section of the tree before it reads anything. Found in review of #1733. --- .../mace-core/src/mace_core/config/base.py | 23 ++++++++- .../mace-core/tests/test_mace_core_config.py | 47 +++++++++++++++++++ 2 files changed, 69 insertions(+), 1 deletion(-) diff --git a/packages/mace-core/src/mace_core/config/base.py b/packages/mace-core/src/mace_core/config/base.py index 8aa5b3d27..00eda70f5 100644 --- a/packages/mace-core/src/mace_core/config/base.py +++ b/packages/mace-core/src/mace_core/config/base.py @@ -237,6 +237,20 @@ def _contains_a_set(annotation: Any) -> bool: ) +def _complete(model: type[BaseModel]) -> None: + """Resolve the tree's forward references and check every section. The + rebuild of an outer model does not rebuild an inner one, so each section + is rebuilt where it is met; the check runs on every section, not only the + rebuilt ones, since a rebuild elsewhere (pydantic's own on first use) + completes a class without checking it.""" + if not model.__pydantic_complete__: + model.model_rebuild() + _check_schema(model) + for field in model.model_fields.values(): + for section, _ in _sections_in(field.annotation): + _complete(section) + + def _check_schema(model: type[BaseModel]) -> None: """Fail at class definition for a field shape the contract cannot keep. @@ -257,6 +271,11 @@ def _check_schema(model: type[BaseModel]) -> None: factory (its kind is the default kind, read without running anything); every section is a `ConfigSection` (a plain `BaseModel` ignores unknown keys, so a typo would vanish, and skips these checks). + + A class whose annotations still name a class defined later, a forward + reference, is incomplete at definition and is not checked here: its + fields cannot be seen through. `_complete` checks it when `load` first + resolves it. """ def reject(name: str, reason: str) -> None: @@ -368,7 +387,8 @@ class ConfigSection(BaseModel): @classmethod def __pydantic_init_subclass__(cls, **kwargs: Any) -> None: super().__pydantic_init_subclass__(**kwargs) - _check_schema(cls) + if cls.__pydantic_complete__: # else a forward reference; `load` checks + _check_schema(cls) @model_validator(mode="before") @classmethod @@ -430,6 +450,7 @@ def load( or an unparsable override, and pydantic's `ValidationError` for a value of the wrong type. """ + _complete(cls) values: dict[str, Any] = {} if config_file is not None: values = read_config_file(config_file) diff --git a/packages/mace-core/tests/test_mace_core_config.py b/packages/mace-core/tests/test_mace_core_config.py index 70c02c124..57e17d1e6 100644 --- a/packages/mace-core/tests/test_mace_core_config.py +++ b/packages/mace-core/tests/test_mace_core_config.py @@ -431,6 +431,53 @@ def double(self) -> int: return 2 * self.seed +# A class that names a class defined below it is incomplete at definition: +# pydantic keeps the name, so the check cannot see through the field. The +# first `load` resolves the name and checks the class then. + + +class Forward(ConfigSection): + later: "Later | None" = None + + +class Later(ConfigSection): + x: int = 1 + + +class ForwardConfig(ReforgeBaseConfig): + forward: Forward = Field(default_factory=Forward) + + +class Leaking(ConfigSection): + plain: "PlainLater" = Field(default_factory=lambda: PlainLater()) + + +class PlainLater(BaseModel): # not a ConfigSection: it would swallow a typo + a: int = 1 + + +class LeakingConfig(ReforgeBaseConfig): + leaking: Leaking = Field(default_factory=Leaking) + + +def test_a_forward_reference_is_checked_and_walked_once_it_resolves(tmp_path): + assert not Forward.__pydantic_complete__ + assert ForwardConfig.load().forward.later is None + config = ForwardConfig.load(cli_overrides=["--forward.later.x", "2"]) + assert config.forward.later == Later(x=2) + path = tmp_path / "typo.json" + path.write_text(json.dumps({"forward": {"later": {"x": 2, "typo": 1}}})) + with pytest.raises(ConfigError, match=r"unknown config key 'forward.later.typo'"): + ForwardConfig.load(path) + + +def test_a_lenient_section_behind_a_forward_reference_is_rejected(): + with pytest.raises( + TypeError, match=r"Leaking.plain holds PlainLater, which is not a ConfigSection" + ): + LeakingConfig.load() + + def test_user_dict_holds_only_what_was_set(tmp_path): config = DemoConfig.load( write_config(tmp_path, ".json"), ["--model.num_interactions", "3"] From 970962c4aff82e484b89b0f8284a759a389df515 Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Tue, 22 Sep 2026 12:59:38 +0200 Subject: [PATCH 05/14] Pin the revision-4 config contract in the tests before the base rewrite (CORE-2, #1556) The tests now pin what `mace_core.config.base` must do after its clean-room rewrite: any number of files in order, settings of every kind kept across layers with the selection separate (`kind:`, `--loss.kind`, bare name), a warning for every override without effect and never for a file, `kind` as the tag only under a kinds field, sections never optional, and unknown keys reported under their file or override. They fail against the module on the branch on purpose; the rewrite follows in the next commit. The kinds error is pinned as one sentence naming the fix and the source of each kind key; the same override twice warns for the first; a variant subclass instance exports under its kind; `.YAML` loads and `load(None)` is `load(())`. Both config modules turn warnings into errors (`pytestmark`), so a spurious `ConfigWarning` fails the suite. --- .../mace-core/tests/test_mace_core_config.py | 271 ++++++++-- .../tests/test_mace_core_config_kinds.py | 499 +++++++++++------- 2 files changed, 554 insertions(+), 216 deletions(-) diff --git a/packages/mace-core/tests/test_mace_core_config.py b/packages/mace-core/tests/test_mace_core_config.py index 57e17d1e6..35e913c10 100644 --- a/packages/mace-core/tests/test_mace_core_config.py +++ b/packages/mace-core/tests/test_mace_core_config.py @@ -1,15 +1,25 @@ """`ReforgeBaseConfig`: file formats, precedence, dotted overrides, unknown keys, -and the resolved export's fixed point.""" +overrides without effect, and the resolved export's fixed point.""" import json +import re import subprocess import sys -from typing import Annotated +import warnings +from typing import Annotated, Any import pytest import yaml -from mace_core.config import ConfigError, ConfigSection, ReforgeBaseConfig -from pydantic import BaseModel, Field, ValidationError, computed_field +from mace_core.config import ( + ConfigError, + ConfigSection, + ConfigWarning, + ReforgeBaseConfig, +) +from pydantic import BaseModel, ConfigDict, Field, ValidationError, computed_field + +#: A warning the test did not ask for is a failure. +pytestmark = pytest.mark.filterwarnings("error") # --------------------------------------------------------------------------- # The demo schema: two levels of nesting, a list, an optional, a Literal. @@ -44,8 +54,8 @@ class DemoConfig(ReforgeBaseConfig): default_dtype: str = "float64" model: ModelSection = ModelSection() data: DataSection = DataSection() - #: An optional section: absent unless the file or the CLI opens it. - stage_two: StageTwoSection | None = None + #: A section left at its defaults unless a file or the CLI writes into it. + stage_two: StageTwoSection = StageTwoSection() #: One config, as a dict. Each format test writes it out and loads it back. @@ -77,8 +87,8 @@ def dump(values, extension): return yaml.safe_dump(values) -def write_config(tmp_path, extension, values=FILE_VALUES): - path = tmp_path / f"config{extension}" +def write_config(tmp_path, extension, values=FILE_VALUES, name="config"): + path = tmp_path / f"{name}{extension}" path.write_text(dump(values, extension), encoding="utf-8") return path @@ -96,6 +106,12 @@ def test_same_config_loads_identically_from_every_format(tmp_path, extension): assert config.model.radial.num_bessel == 8 +def test_extension_is_matched_in_any_case(tmp_path): + path = tmp_path / "CONFIG.YAML" + path.write_text("seed: 5\n", encoding="utf-8") + assert DemoConfig.load(path).seed == 5 + + def test_unknown_extension_is_an_error(tmp_path): path = tmp_path / "config.ini" path.write_text("seed = 1", encoding="utf-8") @@ -109,6 +125,12 @@ def test_empty_file_is_all_defaults(tmp_path): assert DemoConfig.load(path) == DemoConfig() +def test_comment_only_file_is_all_defaults(tmp_path): + path = tmp_path / "comments.yaml" + path.write_text("# nothing set yet\n", encoding="utf-8") + assert DemoConfig.load(path) == DemoConfig() + + def test_file_must_be_a_table_at_the_top(tmp_path): path = tmp_path / "list.json" path.write_text("[1, 2]", encoding="utf-8") @@ -121,6 +143,13 @@ def test_missing_file_is_a_config_error(tmp_path): DemoConfig.load(tmp_path / "nope.yaml") +def test_unreadable_file_is_a_config_error(tmp_path): + path = tmp_path / "latin.yaml" + path.write_bytes(b"name: caf\xe9\n") + with pytest.raises(ConfigError, match=r"cannot read config file .*latin\.yaml"): + DemoConfig.load(path) + + @pytest.mark.parametrize( ("extension", "text"), [(".toml", "seed = \n"), (".yaml", "seed: [1\n"), (".json", "{")], @@ -133,7 +162,7 @@ def test_malformed_file_is_a_config_error(tmp_path, extension, text): # --------------------------------------------------------------------------- -# Precedence: defaults < file < CLI. The legacy behaviour this pins is +# Precedence: defaults < files in order < CLI. The legacy behaviour this pins is # tests/unit/test_arg_parser.py::test_cli_flag_overrides_yaml_config. @@ -141,6 +170,7 @@ def test_no_inputs_gives_the_defaults(): config = DemoConfig.load() assert config == DemoConfig() assert config.model.num_interactions == 2 + assert DemoConfig.load(None) == DemoConfig() # an optional path, unset def test_file_overrides_defaults(tmp_path): @@ -160,6 +190,19 @@ def test_cli_overrides_file_which_overrides_defaults(tmp_path): assert config.default_dtype == "float64" +def test_files_apply_in_order_before_the_overrides(tmp_path): + first = write_config(tmp_path, ".yaml", name="defaults") + second = write_config( + tmp_path, ".toml", {"seed": 8, "model": {"radial": {"num_bessel": 6}}}, "site" + ) + config = DemoConfig.load([first, second], ["--model.num_interactions", "3"]) + assert config.seed == 8 # the second file beats the first + assert config.name == "water" # the first file's other values survive + assert config.model.radial == RadialSection(num_bessel=6, cutoff=4.5) + assert config.model.num_interactions == 3 # the CLI beats both + assert DemoConfig.load([]) == DemoConfig() + + # --------------------------------------------------------------------------- # Dotted CLI overrides @@ -170,10 +213,10 @@ def test_dotted_override_reaches_a_two_level_nested_field(): assert config.model.radial.num_bessel == 8 -def test_dotted_override_opens_an_optional_section(): +def test_dotted_override_reaches_a_section_left_at_its_defaults(): config = DemoConfig.load(cli_overrides=["--stage_two.start_epoch", "50"]) assert config.stage_two == StageTwoSection(start_epoch=50) - assert DemoConfig.load().stage_two is None + assert DemoConfig.load().stage_two == StageTwoSection() def test_override_forms_and_types(): @@ -213,14 +256,25 @@ def test_value_starting_with_dashes_works_in_both_forms(): assert DemoConfig.load(cli_overrides=["--name", "--odd"]).name == "--odd" -def test_dict_valued_field_takes_json_and_is_not_dotted_into(): +@pytest.mark.parametrize("token", ["--", "--=5"]) +def test_bare_dashes_are_an_unknown_option_not_a_key(token): + with pytest.raises(ConfigError, match=rf"unknown config option '{token}'"): + DemoConfig.load(cli_overrides=[token, "--seed", "5"]) + + +def test_overrides_given_as_one_string_are_refused(): + with pytest.raises(TypeError, match="cli_overrides is a string"): + DemoConfig.load(cli_overrides="--seed 5") + + +def test_dict_valued_field_takes_json_and_dotted_paths_into_its_entries(): class Sources(ReforgeBaseConfig): by_name: dict[str, RadialSection] = Field(default_factory=dict) config = Sources.load(cli_overrides=["--by_name", '{"pbe": {"cutoff": 4.0}}']) assert config.by_name == {"pbe": RadialSection(cutoff=4.0)} - with pytest.raises(ConfigError, match=r"unknown config key 'by_name\.pbe\.cutoff'"): - Sources.load(cli_overrides=["--by_name.pbe.cutoff", "4.0"]) + dotted = Sources.load(cli_overrides=["--by_name.pbe.cutoff", "4.0"]) + assert dotted.by_name == {"pbe": RadialSection(cutoff=4.0)} # Inside an entry, the neighbour is still found: the key passes through. with pytest.raises( ConfigError, @@ -229,6 +283,27 @@ class Sources(ReforgeBaseConfig): Sources.load(cli_overrides=["--by_name", '{"pbe": {"cutof": 4.0}}']) +def test_collections_behind_none_or_annotated_keep_their_dotted_paths(tmp_path): + class Collections(ReforgeBaseConfig): + counts: dict[str, int] | None = None + documented: dict[str, Annotated[RadialSection, Field(description="d")]] = Field( + default_factory=dict + ) + pair: tuple[int, RadialSection] | None = None + + assert Collections.load(cli_overrides=["--counts.x", "1"]).counts == {"x": 1} + with pytest.raises( + ConfigError, + match=r"'documented\.a\.cutof'; did you mean 'documented\.a\.cutoff'", + ): + Collections.load(cli_overrides=["--documented", '{"a": {"cutof": 4.0}}']) + path = write_config(tmp_path, ".json", {"pair": [1, {"cutof": 4.0}]}) + with pytest.raises( + ConfigError, match=r"'pair\.1\.cutof'; did you mean 'pair\.1\.cutoff'" + ): + Collections.load(path) + + def test_dict_override_merges_entries_but_list_override_replaces(tmp_path): class Sources(ReforgeBaseConfig): by_name: dict[str, RadialSection] = Field(default_factory=dict) @@ -245,13 +320,11 @@ class Sources(ReforgeBaseConfig): def test_overrides_apply_in_order_on_top_of_the_file(tmp_path): - # Closing a section with null and reopening it drops what the file set - # in it; a dotted value followed by the whole section keeps both. + # A dotted value merges into what the file set in the section; a dotted + # value followed by the whole section keeps both. config = DemoConfig.load( write_config(tmp_path, ".yaml", {"stage_two": {"energy_weight": 5.0}}), cli_overrides=[ - "--stage_two", - "null", "--stage_two.start_epoch", "5", "--model.radial.cutoff", @@ -260,10 +333,96 @@ def test_overrides_apply_in_order_on_top_of_the_file(tmp_path): '{"num_interactions": 3}', ], ) - assert config.stage_two == StageTwoSection(start_epoch=5) + assert config.stage_two == StageTwoSection(start_epoch=5, energy_weight=5.0) assert (config.model.num_interactions, config.model.radial.cutoff) == (3, 4.0) +def test_a_yaml_anchor_does_not_share_an_override(tmp_path): + class Two(ReforgeBaseConfig): + a: dict[str, Any] = Field(default_factory=dict) + b: dict[str, Any] = Field(default_factory=dict) + + path = tmp_path / "anchors.yaml" + path.write_text("a: &empty {}\nb: *empty\n", encoding="utf-8") + config = Two.load(path, ["--a.x", "1"]) + assert (config.a, config.b) == ({"x": "1"}, {}) + + +# --------------------------------------------------------------------------- +# An override that had no effect on the config that runs is a warning; a file +# never warns. + + +class Extras(ReforgeBaseConfig): + seed: int = 1 + extra: dict[str, Any] = Field(default_factory=dict) + + +WITHOUT_EFFECT = { + "a later override at the same path": ( + ["--seed", "1", "--seed", "2"], + ["--seed 1 is overridden: seed is 2"], + ), + "a later override above it": ( + ["--extra.a.b", "2", "--extra.a", "5"], + ["--extra.a.b 2 is overridden: extra.a is 5"], + ), + "a later json override above it": ( + ["--extra.a.b", "2", "--extra.a", "[1]"], + ["--extra.a.b 2 is overridden: extra.a is [1]"], + ), + "part of a json override replaced": ( + ["--extra", '{"a": {"b": 1}}', "--extra.a.b", "2"], + ['--extra {"a": {"b": 1}} is overridden: extra.a.b is 2'], + ), + "a later override below a scalar of it": ( + ["--extra.a", "5", "--extra.a.b", "2"], + ["--extra.a 5 is overridden: extra.a.b is 2"], + ), + "the same override twice: the first is overridden": ( + ["--seed", "2", "--seed", "2"], + ["--seed 2 is overridden: seed is 2"], + ), + "the same json override twice: the first is overridden": ( + ["--extra", '{"a": 1}', "--extra", '{"a": 1}'], + ['--extra {"a": 1} is overridden: extra.a is 1'], + ), + "the last of three at one path is named, once": ( + ["--seed", "1", "--seed", "1", "--seed", "3"], + ["--seed 1 is overridden: seed is 3"], + ), + "a plain key named kind is a key, not a selection": ( + ["--extra.kind", "a", "--extra.kind", "b"], + ["--extra.kind a is overridden: extra.kind is b"], + ), + "an empty mapping above a later key is a merge": ( + ["--extra.a.b", "2", "--extra.a", "{}"], + [], + ), + "an empty mapping below a later key is a merge": ( + ["--extra.a", "{}", "--extra.a.b", "2"], + [], + ), +} + + +@pytest.mark.parametrize("row", WITHOUT_EFFECT, ids=WITHOUT_EFFECT) +def test_an_override_without_effect_warns(row): + tokens, expected = WITHOUT_EFFECT[row] + with warnings.catch_warnings(record=True) as caught: + warnings.simplefilter("always") + Extras.load(cli_overrides=tokens) + assert [str(w.message) for w in caught] == expected + assert all(issubclass(w.category, ConfigWarning) for w in caught) + + +def test_a_file_value_the_cli_replaces_does_not_warn(tmp_path): + with warnings.catch_warnings(): + warnings.simplefilter("error", ConfigWarning) + config = DemoConfig.load(write_config(tmp_path, ".yaml"), ["--seed", "9"]) + assert config.seed == 9 + + # --------------------------------------------------------------------------- # Unknown keys name the key and its nearest neighbour, in files and on the CLI. @@ -294,6 +453,21 @@ def test_every_unknown_key_is_reported_at_once(tmp_path): assert "'model.num_interaction'" in str(excinfo.value) +def test_unknown_keys_are_reported_under_their_file_or_override(tmp_path): + first = write_config(tmp_path, ".yaml", {"sead": 1}, "first") + second = write_config( + tmp_path, ".json", {"model": {"num_interaction": 3}}, "second" + ) + with pytest.raises(ConfigError) as excinfo: + DemoConfig.load([first, second], ["--nmae", "x"]) + assert str(excinfo.value).splitlines() == [ + f"{first}: unknown config key 'sead'; did you mean 'seed'?", + f"{second}: unknown config key 'model.num_interaction'; " + "did you mean 'model.num_interactions'?", + "--nmae x: unknown config key 'nmae'; did you mean 'name'?", + ] + + def test_unknown_key_without_a_close_neighbour_still_names_it(tmp_path): path = tmp_path / "far.yaml" path.write_text("zzzzzz: 1\n", encoding="utf-8") @@ -309,7 +483,7 @@ def test_unknown_dotted_override_names_the_neighbour(): DemoConfig.load(cli_overrides=["--model.num_interaction", "3"]) -def test_unknown_key_inside_an_optional_section(): +def test_unknown_key_inside_a_nested_section(): with pytest.raises( ConfigError, match=r"'stage_two\.start'; did you mean 'stage_two\.start_epoch'\?", @@ -317,15 +491,17 @@ def test_unknown_key_inside_an_optional_section(): DemoConfig.load(cli_overrides=["--stage_two.start", "50"]) -def test_unknown_key_under_a_section_or_scalar_field_drops_the_tag(): - class SectionOrInt(ReforgeBaseConfig): - radial: RadialSection | int = 3 +def test_every_bad_list_item_is_reported(tmp_path): + class Layers(ReforgeBaseConfig): + layers: list[RadialSection] = Field(default_factory=list) - # pydantic tags the location with the member's class name; not a key. - with pytest.raises( - ConfigError, match=r"'radial\.cutof'; did you mean 'radial\.cutoff'\?" - ): - SectionOrInt.load(cli_overrides=["--radial", '{"cutof": 4.0}']) + path = write_config(tmp_path, ".json", {"layers": [{"cutof": 1}, {"nb": 2}]}) + with pytest.raises(ConfigError) as excinfo: + Layers.load(path) + assert re.findall(r"unknown config key '([^']*)'", str(excinfo.value)) == [ + "layers.0.cutof", + "layers.1.nb", + ] def test_help_flag_is_an_error_not_an_exit(): @@ -365,7 +541,7 @@ def test_resolved_dict_has_every_default_in_declaration_order(tmp_path): "data", "stage_two", ] - assert resolved["stage_two"] is None + assert resolved["stage_two"] == {"start_epoch": 100, "energy_weight": 1000.0} assert list(resolved["model"]) == ["num_interactions", "hidden_irreps", "radial"] assert resolved["model"]["radial"] == {"num_bessel": 8, "cutoff": 5.0} assert resolved["data"]["train_file"] is None @@ -382,15 +558,16 @@ def assert_fixed_point(tmp_path, first, extension): @pytest.mark.parametrize("extension", [".yaml", ".json"]) def test_file_to_resolved_to_file_to_resolved_is_a_fixed_point(tmp_path, extension): first = DemoConfig.load( - write_config(tmp_path, ".toml"), ["--model.num_interactions", "3"] + write_config(tmp_path, ".toml"), + ["--model.num_interactions", "3", "--data.train_file", "null"], ).to_resolved_dict() - assert first["stage_two"] is None # a None is part of what has to survive + assert first["data"]["train_file"] is None # a None is part of what has to survive assert_fixed_point(tmp_path, first, extension) def test_fixed_point_holds_through_toml_when_nothing_is_none(tmp_path): - # TOML has no null, so the optional section is opened and the optional - # file name set; the resolved dict then goes through all three formats. + # TOML has no null, so the file sets the optional file name; the resolved + # dict then goes through all three formats. first = DemoConfig.load( write_config(tmp_path, ".yaml"), ["--stage_two.start_epoch", "50"] ).to_resolved_dict() @@ -405,17 +582,25 @@ class LenientSection(BaseModel): def test_field_shapes_the_contract_cannot_keep_are_rejected_at_class_definition(): # Each shape would break a guarantee: set order varies with the hash - # seed; aliases and computed fields do not validate back; a union of - # sections would let a value pick its section; a lenient section would - # swallow typos. + # seed; aliases, excluded and computed fields do not validate back; a + # union of sections would let a value pick its section; a lenient section + # would swallow typos; a section is never optional. shapes = { r"tags is typed as a set.*Use a list": ("tags", list[set[str]]), r"num has an alias": ("num", Annotated[int, Field(alias="n")]), + r"vnum has an alias": ("vnum", Annotated[int, Field(validation_alias="n")]), + r"snum has an alias": ("snum", Annotated[int, Field(serialization_alias="n")]), + r"hidden is excluded from dumps": ( + "hidden", + Annotated[int, Field(exclude=True)], + ), r"either is a union of sections": ("either", RadialSection | StageTwoSection), r"radial holds LenientSection, which is not a ConfigSection": ( "radial", LenientSection | None, ), + r"radial mixes its kinds with int": ("radial", RadialSection | int), + r"stage admits None": ("stage", StageTwoSection | None), } for message, (name, annotation) in shapes.items(): with pytest.raises(TypeError, match=message): @@ -431,13 +616,21 @@ def double(self) -> int: return 2 * self.seed +def test_a_section_cannot_reopen_extra(): + with pytest.raises(TypeError, match=r"Loose sets extra='allow'; a section keeps"): + + class Loose(ConfigSection): + model_config = ConfigDict(extra="allow") + seed: int = 1 + + # A class that names a class defined below it is incomplete at definition: # pydantic keeps the name, so the check cannot see through the field. The # first `load` resolves the name and checks the class then. class Forward(ConfigSection): - later: "Later | None" = None + later: "Later" = Field(default_factory=lambda: Later()) class Later(ConfigSection): @@ -462,7 +655,7 @@ class LeakingConfig(ReforgeBaseConfig): def test_a_forward_reference_is_checked_and_walked_once_it_resolves(tmp_path): assert not Forward.__pydantic_complete__ - assert ForwardConfig.load().forward.later is None + assert ForwardConfig.load().forward.later == Later() config = ForwardConfig.load(cli_overrides=["--forward.later.x", "2"]) assert config.forward.later == Later(x=2) path = tmp_path / "typo.json" @@ -491,7 +684,7 @@ def test_user_dict_holds_only_what_was_set(tmp_path): # --------------------------------------------------------------------------- -# Nothing but the file and the CLI feeds a config. +# Nothing but the files and the CLI feed a config. def test_environment_variables_are_ignored(monkeypatch): diff --git a/packages/mace-core/tests/test_mace_core_config_kinds.py b/packages/mace-core/tests/test_mace_core_config_kinds.py index 2e3d7a24e..63301c20d 100644 --- a/packages/mace-core/tests/test_mace_core_config_kinds.py +++ b/packages/mace-core/tests/test_mace_core_config_kinds.py @@ -1,9 +1,12 @@ """A field of several kinds of section (a discriminated union) under the file -and dotted-override contract. A config writes the kind as the key the section -sits under (`loss: {huber: {delta: 0.1}}`, `--loss.huber.delta 0.1`) or as a -bare name for the kind with its defaults; code sees the union. The file and -every override merge in order, the kind written last wins, the others are -dropped with a warning, two kinds in one place is an error.""" +and dotted-override contract. A config writes each kind's settings under the +kind's name (`loss: {huber: {delta: 0.1}}`, `--loss.huber.delta 0.1`); the +settings of every kind are kept across files and overrides. Which kind runs is +selected by `kind: huber`, `--loss.kind huber` or the bare name `--loss huber`; +a single kind key selects itself; the last selection wins. An override that +changed nothing about the config that runs warns, a file never warns; two kinds +with no selection is an error. `kind` is the tag only under a kinds field: a +plain section field of the same class keeps it as a key. Code sees the union.""" import json import re @@ -19,6 +22,9 @@ ) from pydantic import BaseModel, Field, ValidationError +#: A warning the test did not ask for is a failure. +pytestmark = pytest.mark.filterwarnings("error") + # --------------------------------------------------------------------------- # The schema: a loss of three kinds, one of which holds a field of two kinds; # the same choice with "none of them" as a fourth kind and the default, and @@ -122,8 +128,20 @@ def error(*fragments): return Raises(ConfigError, *fragments) -def ignored(loser, source, winner_source, winner): - return f"{loser} from {source} is ignored: {winner_source} selects {winner}" +def no_effect(source, field, kind, running): + return f"{source}: {field}.{kind} has no effect, {field} runs {running}" + + +def overridden(source, field, running): + return f"{source} is overridden: {field} runs {running}" + + +def needs_kind(field, *kinds): + """The kinds error up to the sources; the first kind is the example.""" + choices = " or ".join(f"kind: {kind}" for kind in kinds) + return ( + f"{field} needs a kind; write {choices} in a file, or pass --{field} {kinds[0]}" + ) def huber(delta=0.01, sub=SubX, **sub_fields): @@ -172,45 +190,34 @@ def layer(index, kind, **fields): SWITCHING = { "file kind, no cli": (HUBER_FILE, "", huber(0.5)), - "cli name switches kind, file section dropped with a warning": ( + "cli name selects; the file's settings of the other kind stay unused": ( HUBER_FILE, "--choice weighted", - Warns( - weighted(), - ignored( - "choice.huber", - "the config file", - "--choice weighted", - "choice.weighted", - ), - ), + weighted(), ), - "cli dotted key switches kind": ( + "cli key of a second kind without a selection is an error": ( HUBER_FILE, "--choice.weighted.stress_weight 5", - Warns( - weighted(5.0), - ignored( - "choice.huber", - "the config file", - "--choice.weighted.stress_weight 5", - "choice.weighted", - ), + error( + needs_kind("choice", "huber", "weighted"), + "; huber from ", + "config.json, weighted from --choice.weighted.stress_weight 5", ), ), - "cli json switches kind": ( + "cli json of a second kind without a selection is an error": ( HUBER_FILE, '--choice {"universal": {"huber_delta": 3}}', - Warns( - universal(3.0), - ignored( - "choice.huber", - "the config file", - '--choice {"universal": {"huber_delta": 3}}', - "choice.universal", - ), + error( + needs_kind("choice", "huber", "universal"), + 'config.json, universal from --choice {"universal": {"huber_delta": 3}}', ), ), + "the tag alone selects on the cli": ({}, '--choice {"kind": "huber"}', huber()), + "the tag beside the settings selects in the file": ( + {"choice": {"kind": "weighted", "huber": {"delta": 0.5}, "weighted": {}}}, + "", + weighted(), + ), "cli name of the file's kind keeps the file's keys": ( HUBER_FILE, "--choice huber", @@ -221,29 +228,30 @@ def layer(index, kind, **fields): "--choice.huber.sub y", huber(0.5, SubY), ), - "switch away and back keeps the file's keys": ( + "a selection, then a key of another kind: the key warns": ( HUBER_FILE, "--choice weighted --choice.huber.sub y", Warns( - huber(0.5, SubY), - ignored( - "choice.weighted", - "--choice weighted", - "--choice.huber.sub y", - "choice.huber", - ), + weighted(), + no_effect("--choice.huber.sub y", "choice", "huber", "weighted"), ), ), - "two kinds on the cli, the later wins": ( + "keys of two kinds on the cli without a selection is an error": ( {}, "--choice.huber.delta 2 --choice.weighted.stress_weight 5", + error( + needs_kind("choice", "huber", "weighted") + + "; huber from --choice.huber.delta 2, " + "weighted from --choice.weighted.stress_weight 5" + ), + ), + "keys of two kinds on the cli, then a selection": ( + {}, + "--choice.huber.delta 2 --choice.weighted.stress_weight 5 --choice huber", Warns( - weighted(5.0), - ignored( - "choice.huber", - "--choice.huber.delta 2", - "--choice.weighted.stress_weight 5", - "choice.weighted", + huber(2.0), + no_effect( + "--choice.weighted.stress_weight 5", "choice", "weighted", "huber" ), ), ), @@ -261,38 +269,21 @@ def layer(index, kind, **fields): "--choice.huber.delta 2", huber(2.0), ), - "bare name in the file, cli other kind": ( + "bare name in the file, cli key of another kind warns": ( {"choice": "huber"}, "--choice.weighted.stress_weight 5", Warns( - weighted(5.0), - ignored( - "choice.huber", - "the config file", - "--choice.weighted.stress_weight 5", - "choice.weighted", + huber(), + no_effect( + "--choice.weighted.stress_weight 5", "choice", "weighted", "huber" ), ), ), "= form": (HUBER_FILE, "--choice.huber.delta=2", huber(2.0)), - "three kinds in a row": ( + "three selections in a row: the last runs, the lost cli one warns": ( HUBER_FILE, "--choice weighted --choice universal", - Warns( - universal(), - ignored( - "choice.huber", - "the config file", - "--choice universal", - "choice.universal", - ), - ignored( - "choice.weighted", - "--choice weighted", - "--choice universal", - "choice.universal", - ), - ), + Warns(universal(), overridden("--choice weighted", "choice", "universal")), ), } @@ -301,14 +292,19 @@ def layer(index, kind, **fields): {"choice": {"huber": {}, "weighted": {}}}, "", error( - "choice is given as several kinds (huber, weighted) in the config " - "file; keep one" + needs_kind("choice", "huber", "weighted"), + "; huber from ", + "config.json, weighted from ", ), ), "two kinds in one json override": ( {}, '--choice {"huber": {}, "weighted": {}}', - error("choice is given as several kinds (huber, weighted) in", "keep one"), + error( + needs_kind("choice", "huber", "weighted") + + '; huber from --choice {"huber": {}, "weighted": {}}, ' + 'weighted from --choice {"huber": {}, "weighted": {}}' + ), ), "unknown kind in the file": ( {"choice": {"hubr": {}}}, @@ -332,44 +328,48 @@ def layer(index, kind, **fields): {}, "--choice.hubr.delta 2", error( - "unknown config key 'choice.hubr.delta'; did you mean 'choice.huber.delta'?" + "unknown config key 'choice.hubr'; did you mean 'choice.huber'?; " + f"the kinds of choice are {KINDS}" ), ), - "key of another kind on the cli is unknown": ( + "a setting beside the kinds is unknown and the kinds are listed": ( {}, "--choice.delta 2", - error("unknown config key 'choice.delta'; did you mean 'choice.huber.delta'?"), + error(f"unknown config key 'choice.delta'; the kinds of choice are {KINDS}"), ), "key of another kind in the file": ( {"choice": {"weighted": {"delta": 2}}}, "", - error( - "unknown config key 'choice.weighted.delta'; did you mean " - "'choice.weighted.stress_weight'?" - ), + error("unknown config key 'choice.weighted.delta'"), ), "the tag is not a key, on the cli": ( {}, "--choice.huber.kind huber", - error("unknown config key 'choice.huber.kind'"), + error( + "unknown config key 'choice.huber.kind'; the key huber already names " + "the kind" + ), ), "the tag is not a key, in the file": ( {"choice": {"huber": {"kind": "weighted"}}}, "", - error("choice.huber.kind is not a key; the kind is given by the key 'huber'"), + error( + "unknown config key 'choice.huber.kind'; the key huber already names " + "the kind" + ), ), - "the tagged form is refused with the key form as the fix": ( + "the flat form fails on the setting beside the tag": ( {"choice": {"kind": "huber", "delta": 2}}, "", - error( - "choice.kind is not a key; write the kind as the key the section " - 'sits under, choice: {"huber": {...}}' - ), + error(f"unknown config key 'choice.delta'; the kinds of choice are {KINDS}"), ), - "the tagged form on the cli": ( + "a selection that is not a kind": ( {}, - '--choice {"kind": "huber"}', - error("choice.kind is not a key"), + "--choice.kind hubr", + error( + "unknown config key 'choice.hubr'; did you mean 'choice.huber'?; " + f"the kinds of choice are {KINDS}" + ), ), "a list at a kinds field": ( {"choice": [1]}, @@ -390,41 +390,36 @@ def layer(index, kind, **fields): "a scalar under a kind": ( {"choice": {"huber": 3}}, "", - error("choice.huber must be a mapping of the kind's keys; got 3"), + error("choice.huber must be a mapping of its keys; got 3"), ), - "null at a kinds field that does not admit it": ( + "null at a kinds field": ( {"choice": None}, "", - error(f"choice does not take null; write a kind, one of {KINDS}"), + error( + "choice must be the name of a kind or a mapping under one, one of " + f"{KINDS}; got null" + ), ), - "a dropped section is not validated": ( + "settings of a kind that does not run are still checked": ( {"choice": {"huber": {"nonsense": 1}}}, "--choice weighted", - Warns( - weighted(), - ignored( - "choice.huber", - "the config file", - "--choice weighted", - "choice.weighted", - ), - ), + error("unknown config key 'choice.huber.nonsense'"), ), } NULL = { - "null override then a key starts the section afresh": ( + "null at a kinds field is an error even when a key follows": ( HUBER_FILE, "--choice null --choice.huber.sub y", - huber(0.01, SubY), + error( + "choice must be the name of a kind or a mapping under one, one of " + f"{KINDS}; got null" + ), ), "null under a kind, on the cli": ( HUBER_FILE, "--choice.huber null", - error( - "choice.huber does not take null; set the keys wanted under it, " - "or write another kind" - ), + error("choice.huber must be a mapping of its keys; got null"), ), "null under an unknown kind is an unknown kind": ( {"choice": {"hubr": None}}, @@ -434,7 +429,7 @@ def layer(index, kind, **fields): "null under a kind, in the file": ( {"choice": {"huber": {"delta": 0.5}, "weighted": None}}, "", - error("choice.weighted does not take null"), + error("choice.weighted must be a mapping of its keys; got null"), ), } @@ -445,21 +440,20 @@ def layer(index, kind, **fields): "none kind, file null": ( {"opt": None}, "", - error(f"opt does not take null; write a kind, one of {OPT_KINDS}"), + error( + "opt must be the name of a kind or a mapping under one, one of " + f"{OPT_KINDS}; got null" + ), ), "none kind, cli null after the file": ( {"opt": {"huber": {}}}, "--opt null", - error(f"opt does not take null; write a kind, one of {OPT_KINDS}"), - ), - "none kind, back to none with a warning": ( - {"opt": {"huber": {}}}, - "--opt none", - Warns( - opt(NoChoice), - ignored("opt.huber", "the config file", "--opt none", "opt.none"), + error( + "opt must be the name of a kind or a mapping under one, one of " + f"{OPT_KINDS}; got null" ), ), + "none kind, back to none": ({"opt": {"huber": {}}}, "--opt none", opt(NoChoice)), } NESTED = { @@ -474,17 +468,17 @@ def layer(index, kind, **fields): "--choice.huber.sub.y.b 7", huber(0.01, SubY, b=7.0), ), - "nested switch warns": ( + "nested selection runs; the file's settings of the other kind stay unused": ( {"choice": {"huber": {"sub": {"y": {"b": 3}}}}}, "--choice.huber.sub x", + huber(0.01, SubX), + ), + "nested key of another kind on the cli warns": ( + {"choice": {"huber": {"sub": {"kind": "y", "y": {"b": 3}}}}}, + "--choice.huber.sub.x.a 5", Warns( - huber(0.01, SubX), - ignored( - "choice.huber.sub.y", - "the config file", - "--choice.huber.sub x", - "choice.huber.sub.x", - ), + huber(0.01, SubY, b=3.0), + no_effect("--choice.huber.sub.x.a 5", "choice.huber.sub", "x", "y"), ), ), "nested empty section is its default kind": ( @@ -492,18 +486,10 @@ def layer(index, kind, **fields): "", huber(0.01, SubX), ), - "outer switch drops the nested section silently": ( + "outer selection: the nested settings stay unused": ( {"choice": {"huber": {"sub": {"y": {"b": 3}}}}}, "--choice weighted", - Warns( - weighted(), - ignored( - "choice.huber", - "the config file", - "--choice weighted", - "choice.weighted", - ), - ), + weighted(), ), "unknown nested kind": ( {"choice": {"huber": {"sub": {"z": {}}}}}, @@ -516,10 +502,7 @@ def layer(index, kind, **fields): "unknown key under a nested kind": ( {"choice": {"huber": {"sub": {"y": {"a": 1}}}}}, "", - error( - "unknown config key 'choice.huber.sub.y.a'; did you mean " - "'choice.huber.sub.y.b'?" - ), + error("unknown config key 'choice.huber.sub.y.a'"), ), } @@ -539,16 +522,21 @@ def layer(index, kind, **fields): '--per_head {"h": {"loss": "huber"}}', per_head("h", Huber, delta=2.0), ), - "kind inside a dict value, json switches with a warning": ( + "kind inside a dict value, json selects another kind": ( {"per_head": {"h": {"loss": {"huber": {"delta": 2}}}}}, '--per_head {"h": {"loss": "weighted"}}', + per_head("h", Weighted), + ), + "kind inside a dict value, dotted key of another kind warns": ( + {"per_head": {"h": {"loss": {"kind": "huber", "huber": {"delta": 2}}}}}, + "--per_head.h.loss.weighted.stress_weight 5", Warns( - per_head("h", Weighted), - ignored( - "per_head.h.loss.huber", - "the config file", - '--per_head {"h": {"loss": "weighted"}}', - "per_head.h.loss.weighted", + per_head("h", Huber, delta=2.0), + no_effect( + "--per_head.h.loss.weighted.stress_weight 5", + "per_head.h.loss", + "weighted", + "huber", ), ), ), @@ -560,10 +548,7 @@ def layer(index, kind, **fields): "unknown key inside a dict value names the kind": ( {"per_head": {"h": {"loss": {"huber": {"stress_weight": 1}}}}}, "", - error( - "unknown config key 'per_head.h.loss.huber.stress_weight'; did you mean " - "'per_head.h.loss.huber.delta'?" - ), + error("unknown config key 'per_head.h.loss.huber.stress_weight'"), ), "kind inside a list item": ( {"layers": [{"loss": "huber"}, {"loss": {"universal": {"huber_delta": 3}}}]}, @@ -578,31 +563,74 @@ def layer(index, kind, **fields): "unknown key inside a list item": ( {"layers": [{"loss": {"huber": {"stress_weight": 1}}}]}, "", + error("unknown config key 'layers.0.loss.huber.stress_weight'"), + ), + "unknown kind inside a list item": ( + {"layers": [{"loss": {"hubr": {}}}]}, + "", error( - "unknown config key 'layers.0.loss.huber.stress_weight'; did you mean " - "'layers.0.loss.huber.delta'?" + "unknown config key 'layers.0.loss.hubr'; did you mean " + f"'layers.0.loss.huber'?; the kinds of layers.0.loss are {KINDS}" ), ), "list item cannot be addressed by a dotted key": ( {}, "--layers.0.loss huber", - error("unknown config key 'layers.0.loss'"), + error("unknown config key 'layers.0'; layers is written whole"), + ), +} + +WARNINGS = { + "a selection lost to a later one warns, the settings stay": ( + {}, + "--choice huber --choice.huber.delta 2 --choice weighted", + Warns( + weighted(), + overridden("--choice huber", "choice", "weighted"), + no_effect("--choice.huber.delta 2", "choice", "huber", "weighted"), + ), + ), + "a selection by the tag lost to a later one warns": ( + {}, + "--choice.kind huber --choice weighted", + Warns(weighted(), overridden("--choice.kind huber", "choice", "weighted")), + ), + "a json override that selects one kind and tunes another warns once": ( + {}, + '--choice {"kind": "weighted", "huber": {"delta": 1, "sub": "y"}}', + Warns( + weighted(), + '--choice {"kind": "weighted", "huber": {"delta": 1, "sub": "y"}}: ' + "choice.huber has no effect, choice runs weighted", + ), + ), + "the same selection twice: the first is overridden": ( + {}, + "--choice huber --choice huber", + Warns(huber(), overridden("--choice huber", "choice", "huber")), + ), + "a setting of the running kind is silent": ( + {"choice": "huber"}, + "--choice.huber.delta 3", + huber(3.0), + ), + "an empty mapping before a setting is silent": ( + {}, + "--choice {} --choice.huber.delta 1", + huber(1.0), + ), + "a file tuning several kinds never warns": ( + {"choice": {"kind": "huber", "huber": {"delta": 0.5}, "weighted": {}}}, + "--energy_weight 2", + lambda c: c.energy_weight == 2.0 and huber(0.5)(c), ), } OTHER = { - "keys of other fields are untouched by a switch": ( + "keys of other fields are untouched by a selection": ( {"energy_weight": 3.0, **HUBER_FILE}, "--choice weighted", - Warns( - lambda c: c.energy_weight == 3.0 and weighted()(c), - ignored( - "choice.huber", - "the config file", - "--choice weighted", - "choice.weighted", - ), - ), + lambda c: c.energy_weight == 3.0 and weighted()(c), ), "a required key of the kind supplied by the cli": ( {"choice": {"huber": {}}}, @@ -616,7 +644,15 @@ def layer(index, kind, **fields): ), } -ROWS = {**SWITCHING, **ERRORS, **NULL, **NONE_KIND, **NESTED, **COLLECTIONS} +ROWS = { + **SWITCHING, + **ERRORS, + **NULL, + **NONE_KIND, + **NESTED, + **COLLECTIONS, + **WARNINGS, +} def split_cli(cli): @@ -743,6 +779,24 @@ def test_user_dict_holds_the_kind_and_only_what_was_set(tmp_path): } +def test_the_empty_mapping_keeps_the_default_instance_with_its_settings(tmp_path): + class Tuned(ReforgeBaseConfig): + choice: Choice = Weighted(stress_weight=2.0) + made: Choice = Field(default_factory=lambda: Huber(delta=9.0, sub=SubY(b=4.0))) + + config = load(tmp_path, Tuned, {"choice": {}, "made": {}}, "") + assert config == Tuned() + assert config.to_user_dict() == { + "choice": {"weighted": {"stress_weight": 2.0}}, + "made": {"huber": {"delta": 9.0, "sub": {"y": {"b": 4.0}}}}, + } + assert load(tmp_path, Tuned, config.to_user_dict(), "") == config + # A kind key with no settings runs the class defaults, not the instance's. + assert load(tmp_path, Tuned, {"choice": {"weighted": {}}}, "").choice == Weighted() + tuned = load(tmp_path, Tuned, {"choice": {}}, "--choice.weighted.stress_weight 3") + assert tuned.choice == Weighted(stress_weight=3.0) + + def test_json_schema_is_produced_in_both_modes(): for mode in ("validation", "serialization"): schema = LossConfig.model_json_schema(mode=mode) @@ -757,6 +811,19 @@ def test_code_sees_the_union_and_may_construct_it_either_way(): assert LossConfig.model_validate(by_tag).choice == Huber(delta=2) +def test_a_variant_class_as_a_plain_field_keeps_kind_as_a_key(tmp_path): + class Reuse(ReforgeBaseConfig): + direct: Huber = Huber() + + resolved = {"direct": {"kind": "huber", "delta": 0.01, "sub": {"x": {"a": 1.0}}}} + assert Reuse().to_resolved_dict() == resolved + assert load(tmp_path, Reuse, resolved, "").to_resolved_dict() == resolved + config = load(tmp_path, Reuse, {}, "--direct.kind huber --direct.sub y") + assert config.direct == Huber(sub=SubY()) + with pytest.raises(ValidationError, match=r"direct\.kind"): + load(tmp_path, Reuse, {}, "--direct.kind weighted") + + # --------------------------------------------------------------------------- # Schema rules. @@ -775,6 +842,13 @@ class TwoTags(ConfigSection): class NamedLikeTheTag(ConfigSection): kind: Literal["kind"] = "kind" + class Colliding(ConfigSection): + """In the flat form `{kind: c, weighted: 9}`, the field would read as + the settings of the sibling kind.""" + + kind: Literal["c"] = "c" + weighted: float = 3.0 + shapes = { r"choice defaults to None, which is not a kind; default to a variant": ( Choice, @@ -812,9 +886,9 @@ class NamedLikeTheTag(ConfigSection): Choice, Plain(), ), - r"choice has a default_factory; write the default as an instance": ( - Weighted | Huber, - Field(discriminator="kind", default_factory=Weighted), + r"choice has variant Colliding with a field named like the kind weighted; ": ( + Annotated[Weighted | Colliding, Field(discriminator="kind")], + Weighted(), ), r"choice mixes its kinds with dict; a kinds field holds its variants only": ( Choice | dict[str, int], @@ -839,11 +913,75 @@ class Required(ReforgeBaseConfig): made: int = Field(default_factory=lambda: factory_calls.append(1) or 1) assert factory_calls == [], "the schema check must not run default factories" - with pytest.raises(ConfigError, match="choice needs a kind; one of"): + with pytest.raises(ConfigError) as excinfo: Required.load(cli_overrides=["--choice", "{}"]) + # Nothing wrote a kind, so no source is named. + assert str(excinfo.value) == ( + "choice needs a kind; write kind: weighted or kind: huber or kind: universal " + "in a file, or pass --choice weighted" + ) assert isinstance(Required.load(cli_overrides=["--choice", "huber"]).choice, Huber) +def test_a_default_factory_that_returns_no_variant_is_reported(tmp_path): + class Broken(ReforgeBaseConfig): + # The wrong result is the point; ty sees only the factory type. + choice: Choice = Field(default_factory=lambda: None) # ty: ignore[invalid-assignment] + + with pytest.raises( + TypeError, match=r"Broken\.choice has a default factory whose result is not" + ): + load(tmp_path, Broken, {"choice": {}}, "") + + +def test_several_kinds_inside_a_list_item_names_only_the_file_fix(tmp_path): + file_values = {"layers": [{"loss": {"huber": {}, "weighted": {}}}]} + with pytest.raises(ConfigError) as excinfo: + load(tmp_path, LossConfig, file_values, "") + path = tmp_path / "config.json" + assert str(excinfo.value) == ( + "layers.0.loss needs a kind; write kind: huber or kind: weighted in a file; " + f"huber from {path}, weighted from {path}" + ) + + +def test_several_files_each_writing_one_kind_are_named(tmp_path): + defaults = tmp_path / "defaults.json" + defaults.write_text(json.dumps({"choice": {"huber": {"delta": 2}}})) + user = tmp_path / "user.json" + user.write_text(json.dumps({"choice": {"weighted": {}}})) + with pytest.raises(ConfigError) as excinfo: + LossConfig.load([defaults, user]) + assert str(excinfo.value) == ( + "choice needs a kind; write kind: huber or kind: weighted in a file, or " + f"pass --choice huber; huber from {defaults}, weighted from {user}" + ) + + +def test_the_tag_inside_a_kind_is_an_error_under_model_validate(): + # Not a silent switch to the other kind: the tag is not a key under a kind. + message = r"choice\.huber\.kind is not a key; huber already names the kind" + with pytest.raises(ValidationError, match=message): + LossConfig.model_validate({"choice": {"huber": {"kind": "weighted"}}}) + with pytest.raises(ValidationError, match=message): + LossConfig.model_validate( + {"choice": {"kind": "huber", "huber": {"kind": "huber", "delta": 2}}} + ) + + +def test_a_subclass_of_a_variant_exports_under_the_variant_kind(tmp_path): + class SubHuber(Huber): + extra_knob: int = 5 + + config = LossConfig(choice=SubHuber(delta=3.0)) + resolved = config.to_resolved_dict() + assert resolved["choice"] == {"huber": {"delta": 3.0, "sub": {"x": {"a": 1.0}}}} + assert config.to_user_dict() == {"choice": {"huber": {"delta": 3.0}}} + reloaded = load(tmp_path, LossConfig, resolved, "") + assert reloaded.choice == Huber(delta=3.0) + assert reloaded.to_resolved_dict() == resolved + + def test_a_kind_named_like_a_class_still_reads_as_a_kind_in_error_paths(): class Root(ConfigSection): kind: Literal["Root"] = "Root" @@ -860,13 +998,20 @@ class Config(ReforgeBaseConfig): def test_a_collection_key_named_like_the_element_class_stays_in_error_paths(): with pytest.raises( - ConfigError, match=r"'heads\.Plain\.q'; did you mean 'heads\.Plain\.p'" + ConfigError, match=r"'heads\.Plain\.pp'; did you mean 'heads\.Plain\.p'" ): - LossConfig.load(cli_overrides=["--heads", '{"Plain": {"q": 1}}']) + LossConfig.load(cli_overrides=["--heads", '{"Plain": {"pp": 1}}']) def test_warning_can_be_turned_into_an_error(tmp_path): with warnings.catch_warnings(): warnings.simplefilter("error", ConfigWarning) - with pytest.raises(ConfigWarning, match=r"choice\.huber from the config file"): - load(tmp_path, LossConfig, HUBER_FILE, "--choice weighted") + with pytest.raises( + ConfigWarning, match=r"choice\.weighted has no effect, choice runs huber" + ): + load( + tmp_path, + LossConfig, + {"choice": "huber"}, + "--choice.weighted.stress_weight 5", + ) From 72417e3789e68b4fac97ca0ba706c9f5976c3944 Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Tue, 22 Sep 2026 12:59:38 +0200 Subject: [PATCH 06/14] Rewrite the config base as one schema walk, one merge and one validation (CORE-2, #1556) `load` flattens every file (any number, in order) and every override into leaf updates carrying their source, walks each path once through the schema (unknown keys and wrong shapes with full dotted paths and neighbours, a bare kind name rewritten to its `kind`, the kinds fields a path enters recorded on the update), merges them set-at-path into one dict, validates once, and then warns for each override that a later override wrote over or that wrote under a kind that does not run (C11): two rules over the override list, the running kind decided from the merged dict by the rule the validator uses. Settings of every kind are kept across layers; the selection is a separate scalar, so a file may tune several kinds and a later layer picks one. `kind` is the tag only by position under a kinds field: a variant class used as a plain section field keeps it as a key on input and export. The kinds error is one `PydanticCustomError` (two kind keys unselected, a required kinds field with nothing written, the tag under a kind), rendered by `load` with the path, the fix and the source of each kind key; it ends the validation of its class, so that class's other errors come on the next load (rebuilding them beside it would need pydantic to know every custom error type). The exports accept a variant subclass instance under the variant's kind. A variant field named like a kind is rejected at class definition (F6). The file extension is case-folded, `load(None)` is `load(())`, and `tomli` is required below Python 3.11 only. The definition-time schema rules (F1-F7, G5) and the annotation introspection move to `config/_schema_rules.py`, which binds `base` as a module so that the import cycle resolves whichever side is imported first. `typing-extensions` is declared for `Self` (the floor is 3.10). --- packages/mace-core/pyproject.toml | 4 +- .../src/mace_core/config/_schema_rules.py | 240 +++ .../mace-core/src/mace_core/config/base.py | 1298 +++++++---------- 3 files changed, 786 insertions(+), 756 deletions(-) create mode 100644 packages/mace-core/src/mace_core/config/_schema_rules.py diff --git a/packages/mace-core/pyproject.toml b/packages/mace-core/pyproject.toml index 487c094a2..60075d184 100644 --- a/packages/mace-core/pyproject.toml +++ b/packages/mace-core/pyproject.toml @@ -22,7 +22,9 @@ dependencies = [ # releases. "pydantic>=2.7", "pyyaml>=6.0", - "tomli>=2.0", + "tomli>=2.0; python_version < '3.11'", + # `typing.Self` is 3.11+; the floor is 3.10. + "typing-extensions>=4.4", ] [project.urls] diff --git a/packages/mace-core/src/mace_core/config/_schema_rules.py b/packages/mace-core/src/mace_core/config/_schema_rules.py new file mode 100644 index 000000000..0358c731a --- /dev/null +++ b/packages/mace-core/src/mace_core/config/_schema_rules.py @@ -0,0 +1,240 @@ +"""The schema rules a config class must keep, checked when the class is defined, +and the annotation introspection the loader shares with them. + +A *section* is a `ConfigSection` subclass; a *kinds field* holds a union of +two or more sections, declared `Annotated[A | B, Field(discriminator="kind")]` +with `kind: Literal["a"]` in each variant. Every rule raises `TypeError` +`. ` naming the fix. + +`base` is bound as a module and dereferenced at call time only: `base.py` +imports this module at its top, so a `from mace_core.config.base import ...` +here would break the package import whichever module is imported first. +""" + +from __future__ import annotations + +import types +from typing import Annotated, Any, Literal, TypeGuard, Union, get_args, get_origin + +from pydantic import BaseModel +from pydantic_core import PydanticUndefined + +from mace_core.config import base as _base + +_NO_KIND_FIX = ( + "default to a variant, or for none of the kinds declare an empty variant " + "with kind: Literal['none']" +) + + +def is_section(node: Any) -> TypeGuard[type[_base.ConfigSection]]: + """A section class. A parametrised generic such as `list[int]` passes + `isinstance(node, type)` on 3.10, hence the origin check first.""" + return ( + get_origin(node) is None + and isinstance(node, type) + and issubclass(node, _base.ConfigSection) + ) + + +def unwrap(node: Any) -> Any: + """`Annotated[X, ...]` as `X`; anything else unchanged.""" + return get_args(node)[0] if get_origin(node) is Annotated else node + + +def members_of(node: Any) -> tuple[Any, ...]: + """The members of a union, nested unions and `Annotated` flattened; a + non-union is its own single member.""" + node = unwrap(node) + if get_origin(node) in (Union, types.UnionType): + return tuple(member for arg in get_args(node) for member in members_of(arg)) + return (node,) + + +def tag_of(variant: type[BaseModel]) -> str | None: + """The one string of a variant's `kind: Literal[...]`; None for any other + shape of tag.""" + field = variant.model_fields.get("kind") + if field is None or get_origin(field.annotation) is not Literal: + return None + values = get_args(field.annotation) + return values[0] if len(values) == 1 and isinstance(values[0], str) else None + + +def kinds_of(node: Any) -> dict[str, type[_base.ConfigSection]] | None: + """`{tag: variant}` for a union of two or more sections, in declaration + order; None for anything else. A variant without a proper tag is keyed by + its class name so that `check_kinds_field` can still name it.""" + members = members_of(node) + if len(members) < 2 or not all(is_section(member) for member in members): + return None + return {tag_of(member) or member.__name__: member for member in members} + + +def kinds_fields_of(cls: type[BaseModel]) -> dict[str, dict[str, type[Any]]]: + """The kinds fields of a class, each with its `{tag: variant}`.""" + fields = {} + for name, field in cls.model_fields.items(): + kinds = kinds_of(field.annotation) + if kinds is not None: + fields[name] = kinds + return fields + + +def reachable_sections(cls: type[BaseModel]) -> list[type[BaseModel]]: + """Every section reachable from `cls` through its fields, `cls` included. + A class left incomplete by a forward reference is rebuilt on the way, which + raises if the name never resolves (F7).""" + found: list[type[BaseModel]] = [] + pending = [cls] + while pending: + section = pending.pop() + if section in found: + continue + if not section.__pydantic_complete__: + section.model_rebuild() + found.append(section) + for field in section.model_fields.values(): + pending.extend(sections_in(field.annotation)) + return found + + +def sections_in(annotation: Any) -> list[type[BaseModel]]: + """The sections inside an annotation: union members and the parameters + of lists, dicts and tuples, at any depth.""" + node = unwrap(annotation) + if is_section(node): + return [node] + return [section for arg in get_args(node) for section in sections_in(arg)] + + +def admits_none(annotation: Any) -> bool: + return any(m in (Any, object, type(None)) for m in members_of(annotation)) + + +def check_model_config(cls: type[BaseModel]) -> None: + """G5: `extra="forbid"` cannot be reopened by a subclass.""" + extra = cls.model_config.get("extra") + if extra != "forbid": + raise TypeError( + f"{cls.__name__} sets extra={extra!r}; a section keeps extra='forbid' " + f"so that an unknown key is an error" + ) + + +def check_fields(cls: type[BaseModel]) -> None: + """F1-F6 over every field of a complete class.""" + for name in cls.model_computed_fields: + raise TypeError( + f"{cls.__name__}.{name} is a computed field; the resolved export " + f"could not be loaded back" + ) + for name, field in cls.model_fields.items(): + where = f"{cls.__name__}.{name}" + if field.alias or field.validation_alias or field.serialization_alias: + raise TypeError( + f"{where} has an alias; the resolved export could not be loaded back" + ) + if field.exclude: + raise TypeError( + f"{where} is excluded from dumps; the resolved export could not be " + f"loaded back" + ) + check_annotation(where, field.annotation, field.discriminator is not None) + kinds = kinds_of(field.annotation) + if kinds is not None: + check_kinds_field(where, kinds, field.default) + elif field.default is None and not admits_none(field.annotation): + raise TypeError(f"{where} defaults to None, which its type does not admit") + + +def _is_lenient_model(node: Any) -> bool: + return ( + get_origin(node) is None + and isinstance(node, type) + and issubclass(node, BaseModel) + and not is_section(node) + ) + + +def check_annotation( + where: str, annotation: Any, discriminated: bool, inside: bool = False +) -> None: + """F1 (no sets), F4 (sections only) and F5 (union shapes) at every depth; + `inside` is true under a dict, list or tuple, where a kinds union is not + allowed.""" + node = unwrap(annotation) + if get_origin(node) is Literal: + return + if node in (set, frozenset) or get_origin(node) in (set, frozenset): + raise TypeError( + f"{where} is typed as a set, whose order changes between runs. Use a list" + ) + members = members_of(node) + for member in members: + if _is_lenient_model(member): + raise TypeError( + f"{where} holds {member.__name__}, which is not a ConfigSection" + ) + sections = [member for member in members if is_section(member)] + if sections: + if type(None) in members: + raise TypeError(f"{where} admits None, which is not a kind; {_NO_KIND_FIX}") + others = [member for member in members if not is_section(member)] + if others: + other = get_origin(others[0]) or others[0] + raise TypeError( + f"{where} mixes its kinds with {other.__name__}; a kinds field holds " + f"its variants only" + ) + if len(sections) > 1 and inside: + raise TypeError( + f"{where} is a union of sections inside a dict, list or tuple; " + f"declare a section holding the kinds field there" + ) + if len(sections) > 1 and not discriminated: + raise TypeError( + f"{where} is a union of sections without a discriminator; declare " + f"kinds as Annotated[A | B, Field(discriminator='kind')]" + ) + return + if len(members) > 1: + for member in members: + check_annotation(where, member, False, inside) + else: + for arg in get_args(node): + check_annotation(where, arg, False, True) + + +def check_kinds_field(where: str, kinds: dict[str, type[Any]], default: Any) -> None: + """F6: one-string tags, no variant named like the tag, no variant field + named like a kind (the flat form could not tell it from a kind's settings), + a variant as the default. A default factory is neither run nor checked here + (F3); its result is checked when an empty mapping selects it.""" + for tag, variant in kinds.items(): + if tag_of(variant) is None: + raise TypeError( + f"{where} has variant {variant.__name__} whose kind must be a " + f"Literal of exactly one string" + ) + if tag == "kind": + raise TypeError( + f"{where} has variant {variant.__name__} whose kind is named 'kind' " + f"like the tag; rename it" + ) + for name in variant.model_fields: + if name != "kind" and name in kinds: + raise TypeError( + f"{where} has variant {variant.__name__} with a field named " + f"like the kind {name}; rename the field" + ) + if default is None: + raise TypeError( + f"{where} defaults to None, which is not a kind; {_NO_KIND_FIX}" + ) + if default is not PydanticUndefined and type(default) not in kinds.values(): + example = next(iter(kinds.values())).__name__ + raise TypeError( + f"{where} has a default that is not one of its variants; write e.g. " + f"{example}()" + ) diff --git a/packages/mace-core/src/mace_core/config/base.py b/packages/mace-core/src/mace_core/config/base.py index 00eda70f5..69ff51cbd 100644 --- a/packages/mace-core/src/mace_core/config/base.py +++ b/packages/mace-core/src/mace_core/config/base.py @@ -1,93 +1,44 @@ -"""The base class every v1 configuration schema derives from. - -A configuration is a tree of fields. A field holds either one value (`seed`, -`cutoff`) or a named group of further fields; such a group is a *section*. -The root of the tree subclasses `ReforgeBaseConfig`, every section -subclasses `ConfigSection`: - - class RadialSection(ConfigSection): - num_bessel: int = 8 - cutoff: float = 5.0 - - class ModelSection(ConfigSection): - num_interactions: int = 2 - radial: RadialSection = RadialSection() - - class TrainConfig(ReforgeBaseConfig): - seed: int = 1 - model: ModelSection = ModelSection() - -Here `model` and `radial` are sections. In a TOML file a section is a table -(`[model.radial]`), in YAML/JSON a nested mapping, and on the command line a -dotted prefix (`--model.radial.cutoff 5.0`). Values come from three layers, -lowest precedence first: - - schema defaults < one config file (.toml/.yaml/.yml/.json) < dotted CLI overrides - -Nothing else feeds a config: no environment variables, no dotenv files, so a -run is reproducible from its file and its command line alone. - -A field may hold a section of one of several *kinds*: a discriminated union -whose tag field names the kind. - - class HuberLoss(ConfigSection): - kind: Literal["huber"] = "huber" - delta: float = 0.01 - - Loss = Annotated[WeightedLoss | HuberLoss, Field(discriminator="kind")] - - class TrainConfig(ReforgeBaseConfig): - loss: Loss = WeightedLoss() - -Code sees the union: `config.loss` is a `WeightedLoss` or a `HuberLoss`. A -file and the command line never write the tag; they write the kind as the -key the section sits under, or as a bare name, which is that kind with -nothing set under it: - - loss: huber --loss huber - loss: {huber: {delta: 0.1}} --loss.huber.delta 0.1 - -The file and each override merge in order. A bare name on top of a section -of the same kind keeps that section's keys; a different kind replaces it: -`--loss weighted` on top of a file with a huber section runs the weighted -loss and warns (`ConfigWarning`) that the huber section is ignored. Two -kinds written in one place, one file or one override, is an error, and so -is `null` at a kinds field or under a kind: a kinds field always holds a -kind. When none of the kinds is a valid choice, that is a kind too, an -empty variant such as `class NoLoss(ConfigSection): kind: Literal["none"]`, -written `loss: none`. A kinds field with no kind written takes the kind of -its default. The resolved and user dicts are written the same way, kind as -key. - -The tagged dict, `{kind: huber, delta: 0.1}`, is pydantic's internal form: -what `model_validate` takes and what `model_json_schema` describes. It is -not a file format; `load()` refuses it. Code builds sections as instances -(`HuberLoss(delta=0.1)`) and never meets either dict form. - -Unknown keys are hard errors. A CLI key is checked against the schema's -dotted paths before anything is built; a key in the file, or inside a -JSON-valued override, is caught by pydantic (`extra="forbid"` on every level -of the tree). Either way the message names the key by its dotted path and, -when there is one, the nearest valid neighbour. - -Field types are restricted to what survives a JSON round trip unchanged, so -that the resolved export is a fixed point: `set` and `frozenset` fields are -rejected when the schema class is defined, because their element order is not -stable across interpreter runs. Use a list. +"""The config base: files and dotted command-line overrides into one validated +pydantic tree, and the tree back out as a dict. + +`load` turns every file (any number, in order) and every override into leaf +updates `(path, value, source)`, walks each path once through the schema (the +one schema-dependent step: unknown keys and wrong shapes are rejected with full +dotted paths, a bare kind name at a kinds field becomes its `kind`, and the +kinds fields the path enters are recorded on the update), merges them +set-at-path into one plain dict, validates that dict once, and finally warns +for each override that lost its effect. Precedence is defaults < files in +order < overrides in order. Nothing else feeds a config: no environment, no +dotenv. + +Wire form of a kinds field (`Annotated[A | B, Field(discriminator="kind")]`): +a mapping with a kind name holding that kind's settings (`huber: {delta: 0.1}`, +any number of them) and `kind: huber` selecting the one that runs; a bare +`huber` is `{kind: huber}`; a single kind key selects itself; `{}` keeps the +schema default. Two or more kind keys without a selection, a required kinds +field with nothing written, and the tag under a kind are the one kinds error, +rendered with the path, the fix and the source of each kind key. The exports +write the kind that ran as `{huber: {fields}}`, without the tag. `kind` is the +tag only under a kinds field: the same class as a plain section field keeps +`kind` as an ordinary key on input and export. + +An override loses its effect in two ways, and only these warn (a file never +does): a later override writes at its path, above it, or below a value of its +that was not a mapping (`{}` never clears a mapping); or it wrote under a kind +that does not run at its kinds field. """ from __future__ import annotations +import copy import difflib import json +import sys import warnings -from collections.abc import Callable, Iterator, Sequence -from itertools import cycle +from collections.abc import Collection, Iterable from pathlib import Path -from types import UnionType -from typing import TYPE_CHECKING, Annotated, Any, Literal, Union, get_args, get_origin +from typing import Any, NamedTuple, TypeVar, get_args, get_origin -import tomli import yaml from pydantic import ( BaseModel, @@ -97,734 +48,571 @@ class TrainConfig(ReforgeBaseConfig): model_serializer, model_validator, ) -from pydantic.fields import FieldInfo - -if TYPE_CHECKING: - from typing_extensions import Self - -__all__ = [ - "ConfigError", - "ConfigSection", - "ConfigWarning", - "ReforgeBaseConfig", - "read_config_file", -] - -#: Config file extensions this module reads, keyed to their parsers. -_FILE_PARSERS = { - ".toml": tomli.loads, +from pydantic_core import ( + ErrorDetails, + PydanticCustomError, + PydanticUndefined, +) +from typing_extensions import Self + +from mace_core.config._schema_rules import ( + check_fields, + check_model_config, + is_section, + kinds_fields_of, + kinds_of, + members_of, + reachable_sections, + unwrap, +) + +if sys.version_info >= (3, 11): + import tomllib +else: + import tomli as tomllib + +_PARSERS = { + ".toml": tomllib.loads, ".yaml": yaml.safe_load, ".yml": yaml.safe_load, ".json": json.loads, } +#: The one error `_to_tagged_form` raises through pydantic; `load` renders it. +_KINDS_ERROR = "kinds" -#: A dotted path as its parts; a list index is a part too. -_Path = tuple[str, ...] - -#: Where a value came from: the layer's index and its name ("the config file" -#: or the override as typed). Comparing two origins compares their order. -_Origin = tuple[int, str] - -#: What runs at a kinds field during a walk: gets the field's value (any -#: shape), the field and its path, returns the value to go on with. -_KindsAction = Callable[[Any, FieldInfo, _Path], Any] - - -# --------------------------------------------------------------------------- -# Schema introspection - - -def _sections_in( - annotation: Any, inside: bool = False -) -> Iterator[tuple[type[BaseModel], bool]]: - """Every section class a field annotation can hold (`Radial`, `Radial | None`, - `list[Radial]`, `Annotated[A | B, ...]`), with whether it sits inside a - dict/list/tuple, where an error location has a key or index before the - section's own field names.""" - origin = get_origin(annotation) # `list` for `list[X]`; None for a plain class - if origin is None: - if isinstance(annotation, type) and issubclass(annotation, BaseModel): - yield annotation, inside - elif origin is Annotated: - yield from _sections_in(get_args(annotation)[0], inside) - elif origin in (Union, UnionType): - for arg in get_args(annotation): - yield from _sections_in(arg, inside) - elif origin in (dict, list, tuple): - for arg in get_args(annotation): - yield from _sections_in(arg, True) - - -def _arms(annotation: Any) -> Iterator[Any]: - """The alternatives of a union, through `Annotated`; else the annotation.""" - origin = get_origin(annotation) - if origin is Annotated: - yield from _arms(get_args(annotation)[0]) - elif origin in (Union, UnionType): - for arg in get_args(annotation): - yield from _arms(arg) - else: - yield annotation - - -def _section_of(annotation: Any) -> tuple[type[BaseModel] | None, bool]: - """The one section a field can hold, and whether it is inside a collection.""" - found = dict(_sections_in(annotation)) - return next(iter(found.items())) if found else (None, False) - - -def _tag_of(field: FieldInfo) -> str | None: - """The tag field name of a kinds field, else None. Also found through an - outer union, `Annotated[...] | None`, which keeps it in the `Annotated` - metadata, so that `_check_schema` can reject that spelling by name.""" - found: Any = field.discriminator - if found is None and get_origin(field.annotation) in (Union, UnionType): - for arg in get_args(field.annotation): - if get_origin(arg) is Annotated: - for meta in get_args(arg)[1:]: - if isinstance(meta, FieldInfo) and meta.discriminator is not None: - found = meta.discriminator - return found if isinstance(found, str) else None - - -def _tag_values(section: type[BaseModel], tag: str) -> tuple[Any, ...]: - """The `Literal` values of a variant's tag field; empty if not a Literal.""" - tag_field = section.model_fields.get(tag) - if tag_field is None or get_origin(tag_field.annotation) is not Literal: - return () - return get_args(tag_field.annotation) - - -def _kinds_of(field: FieldInfo) -> dict[str, type[BaseModel]]: - """Kind name -> section class of a kinds field, in declaration order.""" - tag = _tag_of(field) - if tag is None: - return {} - return { - _tag_values(section, tag)[0]: section - for section, _ in _sections_in(field.annotation) - } - - -def _default_kind(field: FieldInfo) -> str | None: - """The kind of the field's default section; None for a required field.""" - tag = _tag_of(field) - default = field.get_default(call_default_factory=True) - if tag is None or not isinstance(default, BaseModel): - return None - return getattr(default, tag) - - -def _kinds_text(field: FieldInfo) -> str: - return ", ".join(_kinds_of(field)) - - -def _admits_none(annotation: Any) -> bool: - """`X | None`, `Optional[X]`, `Any`, `object`, also under `Annotated`.""" - return any(arm in (type(None), Any, object) for arm in _arms(annotation)) +class ConfigError(ValueError): + """A file or command line the schema cannot take; the message names the + key and the fix.""" -def _contains_a_set(annotation: Any) -> bool: - """A bare `set`, a `set[X]`, or a set anywhere inside, e.g. `list[set[int]]`.""" - if annotation in (set, frozenset) or get_origin(annotation) in (set, frozenset): - return True - return any(_contains_a_set(arg) for arg in get_args(annotation)) +class ConfigWarning(UserWarning): + """An override that had no effect on the config that runs.""" + + +class KindEntered(NamedTuple): + """A kinds field an update's path passes through under one of its kind + keys: the field's path, its kind names and the key entered.""" + + field: tuple[Any, ...] + kinds: tuple[str, ...] + key: str + + +class Update(NamedTuple): + """One leaf update; `source` is the file's path or the override as typed. + The walk fills `under_kinds` with every kinds field the path enters under a + kind key, and `selects` with the kinds field whose tag the path ends at.""" + + path: tuple[Any, ...] + value: Any + source: str + under_kinds: tuple[KindEntered, ...] = () + selects: tuple[Any, ...] | None = None + + +#: The node for the `kind` slot under a kinds field. +_TAG = object() + + +class _KindPosition(NamedTuple): + """A variant class reached through its kinds field, where `kind` is the + tag and not a key. The same class as a plain field is walked as a section.""" + + variant: type[ConfigSection] + + +def _narrowed(node: Any) -> Any: + """`Annotated` stripped and `X | None` stepped through to `X`.""" + members = [m for m in members_of(node) if m is not type(None)] + return unwrap(members[0]) if len(members) == 1 else unwrap(node) + + +def _step(node: Any, key: Any) -> tuple[Any, list[Any] | None]: + """One step of the walk: the child node under `key` (None when the key + is not valid there) and the keys valid at `node` (None when any key is, + `[]` when the node is written whole).""" + node = unwrap(node) + if isinstance(node, _KindPosition): + fields = node.variant.model_fields + keys = [name for name in fields if name != "kind"] + return (fields[key].annotation if key in keys else None), keys + if (kinds := kinds_of(node)) is not None: + keys = ["kind", *kinds] + if key == "kind": + return _TAG, keys + return (_KindPosition(kinds[key]) if key in kinds else None), keys + if is_section(node): + fields = node.model_fields + return (fields[key].annotation if key in fields else None), list(fields) + node = _narrowed(node) + origin = get_origin(node) + if origin is dict: + return get_args(node)[1], None + if node in (Any, object, dict): + return Any, None + if origin in (list, tuple) or node in (list, tuple): + if not isinstance(key, int): + return None, [] + args = get_args(node) + if origin is tuple and Ellipsis not in args and args: + return (args[key] if key < len(args) else Any), [] + return (args[0] if args else Any), [] + return None, [] + + +def _dotted(path: tuple[Any, ...]) -> str: + return ".".join(map(str, path)) + + +def _as_shown(value: Any) -> str: + """JSON where it can, `str` where it cannot (a TOML date, a YAML set).""" + return json.dumps(value, default=str) -_NONE_IS_A_KIND = ( - "default to a variant, or for none of the kinds add an empty variant, " - 'kind: Literal["none"], and default to that' -) +def _as_typed(value: Any) -> str: + """A command-line value as typed (a str), anything else as JSON.""" + return value if isinstance(value, str) else _as_shown(value) -def _complete(model: type[BaseModel]) -> None: - """Resolve the tree's forward references and check every section. The - rebuild of an outer model does not rebuild an inner one, so each section - is rebuilt where it is met; the check runs on every section, not only the - rebuilt ones, since a rebuild elsewhere (pydantic's own on first use) - completes a class without checking it.""" - if not model.__pydantic_complete__: - model.model_rebuild() - _check_schema(model) - for field in model.model_fields.values(): - for section, _ in _sections_in(field.annotation): - _complete(section) - - -def _check_schema(model: type[BaseModel]) -> None: - """Fail at class definition for a field shape the contract cannot keep. - - Each rule protects one guarantee: no sets (order is not stable across - runs, so the export would not be a fixed point); no aliases or computed - fields (the export would not validate back); no `None` default on a type - that does not admit `None` (pydantic does not validate defaults, so the - export would not validate back; a default factory is not run here, so - what it returns is not checked); several sections under one field only - as kinds, i.e. a discriminated union, and not inside a collection (a - value must not become whichever alternative happens to accept it, and - the CLI addresses a collection only as a whole); nothing beside the - variants of a kinds field, not even `None` (a kinds field always holds a - kind; "none of them" is an empty variant, so it can be written, named - among the kinds and warned about like any other); a tag that is one - string other than the tag's own name (it is the key the kind is written - under); a default that is a variant written as an instance, not a - factory (its kind is the default kind, read without running anything); - every section is a `ConfigSection` (a plain `BaseModel` ignores unknown - keys, so a typo would vanish, and skips these checks). - - A class whose annotations still name a class defined later, a forward - reference, is incomplete at definition and is not checked here: its - fields cannot be seen through. `_complete` checks it when `load` first - resolves it. - """ - def reject(name: str, reason: str) -> None: - raise TypeError(f"{model.__name__}.{name} {reason}") - - for name in model.model_computed_fields: - reject(name, "is a computed field; the export must validate back, so drop it") - for name, field in model.model_fields.items(): - if field.alias or field.validation_alias or field.serialization_alias: - reject(name, "has an alias; config keys are field names, so drop it") - if _contains_a_set(field.annotation): - reject( - name, - "is typed as a set; set order is not stable across runs. Use a list", - ) - sections = dict(_sections_in(field.annotation)) - tag = _tag_of(field) - # `default` is undefined, not None, for a required field or a factory; - # a factory is not run here, so what it returns is not checked - if field.default is None and tag is not None: - reject(name, f"defaults to None, which is not a kind; {_NONE_IS_A_KIND}") - if field.default is None and not _admits_none(field.annotation): - reject( - name, "defaults to None but its type does not admit None; add | None" - ) - if len(sections) > 1 and any(sections.values()): - reject( - name, - "is a union of sections inside a dict, list or tuple; the CLI " - "addresses such a field only as a whole. Put the union in a " - "field of the section that is the element", - ) - if len(sections) > 1 and tag is None: - reject( - name, - "is a union of sections without a discriminator; spell it " - 'Annotated[A | B, Field(discriminator="kind")] with a ' - '`kind: Literal["a"]` field in each', - ) - for section in sections: - if not issubclass(section, ConfigSection): - reject( - name, - f"holds {section.__name__}, which is not a ConfigSection; " - f"subclass it", +def _unknown_key( + parent: Any, keys: list[Any] | None, path: tuple[Any, ...] +) -> ConfigError: + """D1: the key by its dotted path, then the nearest neighbour, the kinds + of a kinds field, or why the tag is not a key under a kind.""" + message = f"unknown config key '{_dotted(path)}'" + above = _dotted(path[:-1]) + if keys == []: + return ConfigError( + f"{message}; {above} is written whole and takes no keys under it" + ) + nearest = difflib.get_close_matches( + str(path[-1]), [str(k) for k in keys or []], n=1 + ) + if nearest: + message += f"; did you mean '{_dotted((*path[:-1], nearest[0]))}'?" + parent = unwrap(parent) + if (kinds := kinds_of(parent)) is not None: + message += f"; the kinds of {above} are {', '.join(kinds)}" + elif isinstance(parent, _KindPosition) and path[-1] == "kind": + message += f"; the key {path[-2]} already names the kind" + return ConfigError(message) + + +def _flatten(value: Any, path: tuple[Any, ...], source: str) -> list[Update]: + """Leaf updates of a parsed document: a non-empty mapping recurses, anything + else (a scalar, a whole list, an empty mapping) is a leaf. Each leaf is a + copy, so a YAML anchor shared under two keys becomes two values (G4).""" + if isinstance(value, dict) and value: + return [ + u + for key, item in value.items() + for u in _flatten(item, (*path, key), source) + ] + return [Update(path, copy.deepcopy(value), source)] + + +def _resolve( + root: type[ConfigSection], update: Update, problems: list[str] +) -> Update | None: + """Walk one update's path through the schema. A bare kind name at a kinds + field comes back as its `kind` update; a list value is checked item by item + but stays whole. A `ConfigError` is recorded under the update's source and + the update dropped, so that every problem of one load is reported together.""" + path, value, source = update.path, update.value, update.source + under_kinds: list[KindEntered] = [] + try: + parent, node = None, root + for depth, key in enumerate(path): + child, keys = _step(node, key) + if child is None: + raise _unknown_key(node, keys, path[: depth + 1]) + if isinstance(child, _KindPosition): + kinds = tuple(kinds_of(unwrap(node)) or ()) + under_kinds.append(KindEntered(path[:depth], kinds, key)) + parent, node = node, child + if node is _TAG: + kinds = kinds_of(unwrap(parent)) or {} + if not isinstance(value, str): + raise ConfigError( + f"{_dotted(path)} must name a kind, one of {', '.join(kinds)}; " + f"got {_as_shown(value)}" ) - if tag is None: - continue - others = [arm for arm in _arms(field.annotation) if arm not in sections] - if type(None) in others: - reject(name, f"admits None, which is not a kind; {_NONE_IS_A_KIND}") - if others: - reject( - name, - f"mixes its kinds with {getattr(others[0], '__name__', others[0])}; " - f"a kinds field holds its variants only", - ) - for section in sections: - values = _tag_values(section, tag) - if len(values) != 1 or not isinstance(values[0], str): - reject( - name, - f"has variant {section.__name__} whose {tag} must be a Literal " - f"of exactly one string; a config names the kind by it", + path, node = path[:-1], parent + node = unwrap(node) + if (kinds := kinds_of(node)) is not None: + if isinstance(value, str): + if value not in kinds: + raise _unknown_key(node, ["kind", *kinds], (*path, value)) + return Update((*path, "kind"), value, source, tuple(under_kinds), path) + if value != {}: + raise ConfigError( + f"{_dotted(path)} must be the name of a kind or a mapping under " + f"one, one of {', '.join(kinds)}; got {_as_shown(value)}" ) - if values[0] == tag: - reject( - name, - f"has variant {section.__name__} whose kind is named {tag!r} like " - f"the tag; a config could not tell the two apart. Rename it", + elif isinstance(node, _KindPosition) or is_section(node): + if not isinstance(value, dict): + raise ConfigError( + f"{_dotted(path)} must be a mapping of its keys; " + f"got {_as_shown(value)}" ) - example = f"{next(iter(sections)).__name__}()" - if field.default_factory is not None: - reject( - name, - f"has a default_factory; write the default as an instance, e.g. " - f"{example} (pydantic copies it per instance)", - ) - if not field.is_required() and not isinstance(field.default, tuple(sections)): - reject( - name, - f"has a default that is not one of its variants; write e.g. {example}", - ) + else: + node = _narrowed(node) + is_list = get_origin(node) in (list, tuple) or node in (list, tuple) + if is_list and isinstance(value, list): + for index, item in enumerate(value): + for item_update in _flatten(item, (*path, index), source): + _resolve(root, item_update, problems) + return Update(path, value, source, tuple(under_kinds)) + except ConfigError as error: + problems.append(f"{source}: {error}") + return None -class ConfigError(ValueError): - """A config file or override the schema rejects. - - Raised for a missing, unparsable or malformed file, an unknown key or - kind, an override the CLI parser cannot make sense of, two kinds of one - section written in one place, and `null` at or under a kind. The - message names the file or the - offending key by its dotted path, and suggests the nearest valid key - when there is a close match. +def _resolve_all(root: type[ConfigSection], updates: list[Update]) -> list[Update]: + problems: list[str] = [] + resolved = [_resolve(root, update, problems) for update in updates] + if problems: + raise ConfigError("\n".join(problems)) + return [update for update in resolved if update is not None] + + +def _set_at_path(mapping: dict[Any, Any], path: tuple[Any, ...], value: Any) -> None: + """Set `value` at `path`, creating mappings on the way; an empty mapping + never clears a mapping already there, and is stored as a fresh one so that + later writes under it leave the update's value alone.""" + if not path: + return + for key in path[:-1]: + if not isinstance(mapping.get(key), dict): + mapping[key] = {} + mapping = mapping[key] + if value != {} or not isinstance(mapping.get(path[-1]), dict): + mapping[path[-1]] = {} if value == {} else value + + +def read_config_file(config_file: str | Path) -> dict[str, Any]: + """Parse one file by its extension (`.toml`, `.yaml`, `.yml`, `.json`, in + any case). + + A `ConfigError` names the file for an unknown extension, a file that cannot + be read or parsed, and a top level that is not a table. An empty or + comment-only file is `{}`. Non-string YAML keys (`1:`) stay as parsed; + quote them where the schema wants strings. """ - - -class ConfigWarning(UserWarning): - """A section of one kind is ignored because a later layer selected another - kind. `warnings.simplefilter("error", ConfigWarning)` makes it an error.""" + path = Path(config_file) + parser = _PARSERS.get(path.suffix.lower()) + if parser is None: + raise ConfigError( + f"unknown extension '{path.suffix}' of config file {path}; use .toml, " + f".yaml, .yml or .json" + ) + try: + text = path.read_text(encoding="utf-8") + except (OSError, UnicodeDecodeError) as error: + raise ConfigError(f"cannot read config file {path}: {error}") from error + try: + document = parser(text) + except (ValueError, yaml.YAMLError) as error: + raise ConfigError(f"cannot parse config file {path}: {error}") from error + if document is None: + document = {} + if not isinstance(document, dict): + raise ConfigError( + f"config file {path} must have a table of keys at the top level, not " + f"{type(document).__name__}" + ) + return document + + +def _parse_overrides(tokens: Iterable[str]) -> list[Update]: + """`--a.b.c value` or `--a.b.c=value`, in order. A value that is `null` or + starts with `[` or `{` is JSON (and flattened like a file); any other value + stays a string for pydantic to coerce. Keys match exactly.""" + argv = list(tokens) + updates = [] + position = 0 + while position < len(argv): + token = argv[position] + position += 1 + key, has_inline_value, value = token[2:].partition("=") + if not token.startswith("--") or not key: + raise ConfigError( + f"unknown config option '{token}'; options are --key.path value" + ) + if not has_inline_value: + if position == len(argv): + raise ConfigError(f"override {token} is missing its value") + value = argv[position] + position += 1 + source = token if has_inline_value else f"{token} {value}" + parsed: Any = value + if value == "null" or value[:1] in ("[", "{"): + try: + parsed = json.loads(value) + except ValueError as error: + raise ConfigError( + f"override {token} is not valid JSON: {error}" + ) from error + updates.extend(_flatten(parsed, tuple(key.split(".")), source)) + return updates class ConfigSection(BaseModel): - """A nested section of a configuration: a table in the file, a dotted - prefix on the command line. Unknown keys are errors here too.""" + """A node of a config tree: unknown keys are errors, and the schema rules + of `_schema_rules` are checked when a subclass is defined (or, behind a + forward reference, on the first `load`). + + A kinds field takes the wire form `{kind: X, X: {...}, Y: {...}}`, `{X: + {...}}`, `"X"` or `{}` and dumps as `{X: {fields}}`; internally pydantic + sees the tagged form `{kind: X, ...X's fields}`, which `model_validate` and + direct construction accept as well. An instance of a subclass of a variant + is held as given and exported under the variant's kind. + """ model_config = ConfigDict(extra="forbid") @classmethod def __pydantic_init_subclass__(cls, **kwargs: Any) -> None: super().__pydantic_init_subclass__(**kwargs) - if cls.__pydantic_complete__: # else a forward reference; `load` checks - _check_schema(cls) + check_model_config(cls) + if cls.__pydantic_complete__: + check_fields(cls) @model_validator(mode="before") @classmethod - def _kinds_from_keys(cls, values: Any) -> Any: - """`{huber: {delta: 0.1}}` under a kinds field becomes the - `{kind: huber, delta: 0.1}` pydantic's discriminator reads.""" - if not isinstance(values, dict): - return values - values = dict(values) - for name, field in cls.model_fields.items(): - tag, value = _tag_of(field), values.get(name) - if tag is not None and isinstance(value, dict) and len(value) == 1: - ((kind, inner),) = value.items() - if isinstance(inner, dict) and tag not in value: - values[name] = {**inner, tag: kind} - return values - + def _kinds_fields_to_tagged_form(cls, data: Any) -> Any: + if not isinstance(data, dict): + return data + for name, kinds in kinds_fields_of(cls).items(): + if name in data: + data = {**data, name: _to_tagged_form(cls, name, data[name], kinds)} + return data + + # No return annotation: with one, the serialization JSON schema collapses. @model_serializer(mode="wrap") - def _kinds_as_keys(self, handler: SerializerFunctionWrapHandler): - """The dump with every kinds field written kind-as-key. The kind is read - from the instance: an unset default tag is absent from a user dict. - No return annotation: pydantic would take it as the JSON schema.""" + def _kinds_fields_under_their_name(self, handler: SerializerFunctionWrapHandler): dumped = handler(self) - for name, field in type(self).model_fields.items(): - tag = _tag_of(field) - if tag is not None and isinstance(dumped.get(name), dict): - inner = {key: v for key, v in dumped[name].items() if key != tag} - dumped[name] = {getattr(getattr(self, name), tag): inner} + for name, kinds in kinds_fields_of(type(self)).items(): + if name in dumped: # absent under exclude_unset + variant = getattr(self, name) + if not isinstance(variant, tuple(kinds.values())): + raise TypeError( + f"{type(self).__name__}.{name} holds {type(variant).__name__}, " + f"which is not one of its variants" + ) + settings = {k: v for k, v in dumped[name].items() if k != "kind"} + dumped[name] = {variant.kind: settings} return dumped -class ReforgeBaseConfig(ConfigSection): - """Root of a configuration tree. Subclass it; nest `ConfigSection`s in it. +def _selected_kind(value: dict[str, Any], kinds: Collection[str]) -> str | None: + """The kind a kinds field's wire mapping runs: its `kind` scalar, else its + sole kind key; None when nothing selects one (an empty mapping runs the + schema default, several kind keys need a selection).""" + selected = value.get("kind") + if isinstance(selected, str): + return selected + written = [key for key in value if key in kinds] + return written[0] if len(written) == 1 else None + + +def _kinds_error( + field: str, kinds: list[str], under: str | None = None +) -> PydanticCustomError: + """The kinds error: `field` needs a kind (one of `kinds`), or the tag was + written under the kind `under`. `load` renders it with the path, the fix + and the sources (`_kind_error_message`).""" + choices = " or ".join(f"kind: {kind}" for kind in kinds) + template = ( + "{field} needs a kind; write {choices}" + if under is None + else "{field}.{under}.kind is not a key; {under} already names the kind" + ) + context = {"field": field, "kinds": kinds, "under": under, "choices": choices} + return PydanticCustomError(_KINDS_ERROR, template, context) + + +def _to_tagged_form( + cls: type[ConfigSection], name: str, value: Any, kinds: dict[str, type[Any]] +) -> Any: + """Wire form of one kinds field to pydantic's tagged form. Shapes the + schema cannot take are returned unchanged for pydantic to report; the + shapes it would misreport raise the kinds error, among them a `kind` inside + the selected kind's settings, which is not a key there.""" + if isinstance(value, str): + return {"kind": value} + if not isinstance(value, dict) or not isinstance(value.get("kind", ""), str): + return value + selected = _selected_kind(value, kinds) + if selected is None: + written = [key for key in value if key in kinds] + if written: + raise _kinds_error(name, written) + if value: + return value + default = cls.model_fields[name].get_default(call_default_factory=True) + if default is PydanticUndefined: + raise _kinds_error(name, list(kinds)) + if type(default) not in kinds.values(): + example = next(iter(kinds.values())).__name__ + raise TypeError( + f"{cls.__name__}.{name} has a default factory whose result is not " + f"one of its variants; return e.g. {example}()" + ) + return {"kind": default.kind, **default.model_dump(exclude_unset=True)} + settings = value.get(selected, {}) + if not isinstance(settings, dict): + return value + if "kind" in settings: + raise _kinds_error(name, list(kinds), under=selected) + rest = {k: v for k, v in value.items() if k not in kinds and k != "kind"} + return {**rest, **settings, "kind": selected} - `load()` reads a file and applies overrides. Constructing the class - directly behaves like a plain Pydantic model. - """ + +_ConfigT = TypeVar("_ConfigT", bound="ReforgeBaseConfig") + + +class ReforgeBaseConfig(ConfigSection): + """The root of a config tree: `load` builds it from files and the command + line; the two exports write it back as JSON-native dicts.""" @classmethod def load( cls, - config_file: str | Path | None = None, - cli_overrides: Sequence[str] = (), + config_files: str | Path | Iterable[str | Path] | None = (), + cli_overrides: Iterable[str] = (), ) -> Self: - """Build the config from defaults, then the file, then the overrides. - - `cli_overrides` is the argument list after the program name, e.g. - `["--model.num_interactions", "3", "--seed=7"]`: a dotted path names - a field at any depth. A value starting with `[` or `{`, or the word - `null`, is JSON, so a whole section, a list or a dict can be given; - any other value is a string pydantic converts to the field's type. - An override merges into the file, and into earlier overrides, like a - section does: a dict-valued field gains or replaces entries, so an - entry cannot be removed from the command line; a list-valued field - is replaced whole. Under a kinds field the kind written last wins - and the others are dropped with a `ConfigWarning`. - - Raises `ConfigError` for an unknown key or kind, an unreadable file - or an unparsable override, and pydantic's `ValidationError` for a - value of the wrong type. + """Build the config from `config_files` in order, then `cli_overrides` + in order (`sys.argv[1:]`-style tokens); a lone path is one file, None + is no file. + + Raises `ConfigError` for a file or override the schema cannot take + (every unknown key of the load together, each under its file or + override), pydantic's `ValidationError` for a value of the wrong type, + and warns `ConfigWarning` for each override that had no effect on the + config that runs. """ - _complete(cls) - values: dict[str, Any] = {} - if config_file is not None: - values = read_config_file(config_file) - layers = [(values, "the config file"), *_parse_overrides(cls, cli_overrides)] + if isinstance(cli_overrides, str): + raise TypeError( + "cli_overrides is a string; pass the tokens as a list, " + "like sys.argv[1:]" + ) + for section in reachable_sections(cls): + check_fields(section) + if config_files is None: + config_files = () + files = ( + [config_files] + if isinstance(config_files, (str, Path)) + else list(config_files) + ) + from_files = [ + update + for file in files + for update in _flatten(read_config_file(file), (), str(Path(file))) + ] + updates = _resolve_all(cls, [*from_files, *_parse_overrides(cli_overrides)]) merged: dict[str, Any] = {} - origins: dict[_Path, _Origin] = {} - for index, (layer, source) in enumerate(layers): - layer = _at_kinds(layer, cls, _name_as_mapping) - merged = _deep_update(merged, layer) - _record_origins(origins, layer, (index, source)) - merged = _at_kinds(merged, cls, _KindSelector(origins)) - try: - return cls.model_validate(merged) - except ValidationError as error: - unknown = _unknown_key_messages(cls, error) - if not unknown: - raise - raise ConfigError("\n".join(unknown)) from error + for update in updates: + _set_at_path(merged, update.path, update.value) + config = _validated(cls, merged, updates) + overrides = updates[len(from_files) :] + messages = [ + _override_without_effect(update, overrides[position + 1 :], merged) + for position, update in enumerate(overrides) + ] + for message in dict.fromkeys(m for m in messages if m is not None): + warnings.warn(message, ConfigWarning, stacklevel=2) + return config def to_resolved_dict(self) -> dict[str, Any]: - """Every field, defaults included, as JSON-native values, in schema - order, a kinds field as `{kind: {...}}`. Loading the result back and - resolving again gives the same dict. (TOML has no null, so a `None` - can only go out as YAML or JSON.)""" + """Every field with defaults filled, JSON-native, in declaration order; + loading it back gives an equal config and the same dict. A kinds field + holding a subclass of its variant is written as the variant, so a + field the subclass added is not written.""" return self.model_dump(mode="json") def to_user_dict(self) -> dict[str, Any]: - """Only the fields the file and the overrides set, as JSON-native - values: what the user wrote, for the model metadata.""" + """Only what the files and the overrides set, in the same shape; the + kind that ran is recorded even when it was chosen by default.""" return self.model_dump(mode="json", exclude_unset=True) -def read_config_file(path: str | Path) -> dict[str, Any]: - """Parse one config file, choosing the parser by extension. - - An empty TOML or YAML file is an empty config. Raises `ConfigError` for a missing - file, an unknown extension, a file its parser rejects, or a file whose - top level is not a table. - """ - path = Path(path) - parser = _FILE_PARSERS.get(path.suffix.lower()) - if parser is None: - raise ConfigError( - f"cannot read config file {path}: unknown extension {path.suffix!r}; " - f"expected one of {', '.join(_FILE_PARSERS)}" - ) +def _validated( + cls: type[_ConfigT], merged: dict[str, Any], updates: list[Update] +) -> _ConfigT: + """Validate once; a kinds error comes out as one `ConfigError`, and a load + without a kinds error passes pydantic's `ValidationError` through (D3). A + kinds error ends the validation of its class, so the other errors of that + class are reported on the next load.""" try: - values = parser(path.read_text(encoding="utf-8")) - except OSError as error: - raise ConfigError( - f"cannot read config file {path}: {error.strerror}" - ) from error - except (ValueError, yaml.YAMLError) as error: # tomli/json errors are ValueErrors - raise ConfigError(f"cannot parse config file {path}: {error}") from error - if values is None: - return {} - # TOML always yields a table, but a YAML or JSON file can hold a list or a - # scalar, which would crash the merge with the overrides instead of - # naming the file. - if not isinstance(values, dict): - raise ConfigError( - f"config file {path} must hold a table of keys at the top level, " - f"not a {type(values).__name__}" - ) - return values - - -# --------------------------------------------------------------------------- -# Merging - - -def _deep_update(base: dict[str, Any], update: dict[str, Any]) -> dict[str, Any]: - """`base` overlaid with `update`, recursing where both hold a dict.""" - merged = dict(base) - for key, value in update.items(): - if isinstance(value, dict) and isinstance(merged.get(key), dict): - merged[key] = _deep_update(merged[key], value) - else: - merged[key] = value - return merged - - -def _record_origins( - origins: dict[_Path, _Origin], values: Any, origin: _Origin, path: _Path = () -) -> None: - """Note `origin` for every path a layer writes, sections included.""" - items: Any = () - if isinstance(values, dict): - items = values.items() - elif isinstance(values, list): - items = enumerate(values) - for key, value in items: - origins[(*path, str(key))] = origin - _record_origins(origins, value, origin, (*path, str(key))) - - -def _parse_overrides( - model: type[BaseModel], cli_overrides: Sequence[str] -) -> list[tuple[dict[str, Any], str]]: - """Each `--a.b value` or `--a.b=value` pair as a mapping nested under its - path, with the override as typed, in order.""" - valid = list(_dotted_paths(model)) - valid_set = set(valid) - unknown: list[str] = [] - tokens = iter(cli_overrides) - layers = [] - - for token in tokens: - option = token.removeprefix("--") - if "=" in option: - name, value = option.split("=", 1) - else: - name, value = option, None - if not token.startswith("--") or not name: - raise ConfigError( - f"unknown config option {token!r}; overrides are written --key value" - ) - source = token - if value is None: - value = next(tokens, None) - if value is None: - raise ConfigError(f"override --{name} is missing its value") - source = f"{token} {value}" - - if name not in valid_set: - unknown.append(_unknown_key_message(name, valid)) - continue - - if value == "null" or value.startswith(("[", "{")): - try: - value = json.loads(value) - except ValueError as error: - raise ConfigError( - f"override --{name} is not valid JSON: {error}" - ) from error - - override: Any = value - for section in reversed(name.split(".")): - override = {section: override} - layers.append((override, source)) - - if unknown: - raise ConfigError("\n".join(unknown)) - - return layers - - -def _dotted_paths(model: type[BaseModel], prefix: str = "") -> Iterator[str]: - """Every field of the tree as a dotted path, sections included. A kinds - field lists each kind as a key with the kind's fields under it, without - the tag field: the key names the kind. - - A section inside a dict or list is not descended into: the CLI addresses - such a field only as a whole, with a JSON value. - """ - for name, field in model.model_fields.items(): - path = f"{prefix}{name}" - yield path - kinds = _kinds_of(field) - for kind, variant in kinds.items(): - yield f"{path}.{kind}" - for sub_path in _dotted_paths(variant, f"{path}.{kind}."): - if sub_path != f"{path}.{kind}.{_tag_of(field)}": - yield sub_path - section, inside_collection = _section_of(field.annotation) - if section is not None and not inside_collection and not kinds: - yield from _dotted_paths(section, f"{path}.") - - -# --------------------------------------------------------------------------- -# Kinds: a walk over the values against the schema, with an action at every -# kinds field. The action sees and returns the kind-as-key form. - - -def _at_kinds( - values: Any, section: type[BaseModel] | None, action: _KindsAction, path: _Path = () -) -> Any: - """`values`, a mapping for `section`, with `action` applied at each kinds - field and every section under it walked in turn.""" - if section is None or not isinstance(values, dict): - return values - out = dict(values) - for key, value in values.items(): - field = section.model_fields.get(key) - if field is None: - continue - kinds = _kinds_of(field) - if not kinds: - out[key] = _under(value, field.annotation, action, (*path, key)) - continue - value = action(value, field, (*path, key)) - if isinstance(value, dict): - value = { - kind: _at_kinds(inner, kinds.get(kind), action, (*path, key, kind)) - for kind, inner in value.items() - } - out[key] = value - return out - - -def _under(value: Any, annotation: Any, action: _KindsAction, path: _Path) -> Any: - """`value` with `_at_kinds` applied to every section the annotation reaches - through unions, dicts, lists and tuples. A value whose shape the annotation - does not describe is returned as is, for pydantic to report.""" - origin = get_origin(annotation) - if origin is None: - return _at_kinds(value, _section_of(annotation)[0], action, path) - if origin is Annotated: - return _under(value, get_args(annotation)[0], action, path) - if origin in (Union, UnionType): # at most one arm takes a dict or a list - for arm in get_args(annotation): - value = _under(value, arm, action, path) - return value - if origin is dict and isinstance(value, dict): - value_type = get_args(annotation)[1] - return { - k: _under(v, value_type, action, (*path, str(k))) for k, v in value.items() - } - if origin in (list, tuple) and isinstance(value, list): - item_types = [a for a in get_args(annotation) if a is not Ellipsis] - return [ - _under(v, t, action, (*path, str(i))) - for i, (v, t) in enumerate(zip(value, cycle(item_types), strict=False)) - ] - return value - - -def _shown(value: Any) -> str: - """A value as the user could have written it; a date or a YAML set as text.""" - return json.dumps(value, default=str) - - -def _name_as_mapping(value: Any, field: FieldInfo, path: _Path) -> Any: - """A bare kind name is that kind with its defaults, `{huber: {}}`, so that - it merges into an earlier section of the same kind instead of replacing it.""" - return {value: {}} if isinstance(value, str) else value - - -class _KindSelector: - """At a kinds field after the merge: keep the kind written last, drop the - others with a warning, fill in the default kind, and check the shape.""" - - def __init__(self, origins: dict[_Path, _Origin]) -> None: - self.origins = origins - - def __call__(self, value: Any, field: FieldInfo, path: _Path) -> Any: - dotted = ".".join(path) - tag, kinds, named = _tag_of(field), _kinds_of(field), _kinds_text(field) - if value is None: - raise ConfigError( - f"{dotted} does not take null; write a kind, one of {named}" - ) - if not isinstance(value, dict): - raise ConfigError( - f"{dotted} must be the name of a kind or a mapping under one, " - f"one of {named}; got {_shown(value)}" - ) - if tag in value: - raise ConfigError( - f"{dotted}.{tag} is not a key; write the kind as the key the " - f"section sits under, {dotted}: {{{_shown(value[tag])}: {{...}}}}" - ) - nulled = [k for k, inner in value.items() if inner is None and k in kinds] - if nulled: - raise ConfigError( - f"{dotted}.{nulled[0]} does not take null; set the keys wanted " - f"under it, or write another kind" - ) - if not value: - default = _default_kind(field) - if default is None: - raise ConfigError(f"{dotted} needs a kind; one of {named}") - return {default: {}} - present = value - origin_of = {k: self.origins[(*path, str(k))] for k in present} - by_origin = sorted(present, key=origin_of.__getitem__) - kind = by_origin[-1] - index, source = origin_of[kind] - tied = [k for k in by_origin if origin_of[k][0] == index] - if len(tied) > 1: - raise ConfigError( - f"{dotted} is given as several kinds ({', '.join(tied)}) in " - f"{source}; keep one" - ) - if kind not in kinds: - valid = [f"{dotted}.{k}" for k in kinds] - raise ConfigError( - _unknown_key_message(f"{dotted}.{kind}", valid) - + f"; the kinds of {dotted} are {named}" - ) - for loser in by_origin[:-1]: - warnings.warn( - f"{dotted}.{loser} from {origin_of[loser][1]} is " - f"ignored: {source} selects {dotted}.{kind}", - ConfigWarning, - stacklevel=2, - ) - inner = present[kind] - if not isinstance(inner, dict): - raise ConfigError( - f"{dotted}.{kind} must be a mapping of the kind's keys; " - f"got {_shown(inner)}" - ) - if tag in inner: - raise ConfigError( - f"{dotted}.{kind}.{tag} is not a key; the kind is given by the " - f"key {kind!r}" - ) - return {kind: inner} - - -# --------------------------------------------------------------------------- -# Error messages - - -def _unknown_key_message(key: str, candidates: Sequence[str]) -> str: - message = f"unknown config key {key!r}" - closest = difflib.get_close_matches(key, candidates, n=1) - if closest: - message += f"; did you mean {closest[0]!r}?" + return cls.model_validate(merged) + except ValidationError as error: + kind_errors = [e for e in error.errors() if e["type"] == _KINDS_ERROR] + if not kind_errors: + raise + lines = [_kind_error_message(e, updates) for e in kind_errors] + raise ConfigError("\n".join(lines)) from error + + +def _wrote(update: Update, target: tuple[Any, ...]) -> bool: + """Whether the update wrote at or under `target`, or a whole list holding it.""" + depth = len(update.path) + if depth >= len(target): + return update.path[: len(target)] == target + return update.path == target[:depth] and isinstance(update.value, list) + + +def _kind_error_message(error: ErrorDetails, updates: list[Update]) -> str: + """Pydantic's location is the section that raised, and the field is in the + context. The dotted fix is left out inside a list item, where the walk + would reject it; each kind that was written is named with the source that + first wrote it. The tag under a kind is rejected by the walk before + validation; should it arrive, pydantic's own sentence is kept.""" + ctx = error.get("ctx") or {} + if ctx.get("under") is not None: + return error["msg"] + path = (*error["loc"], ctx["field"]) + kinds = ctx["kinds"] + message = f"{_dotted(path)} needs a kind; write {ctx['choices']} in a file" + if not any(isinstance(part, int) for part in path): + message += f", or pass --{_dotted(path)} {kinds[0]}" + writers = [] + for kind in kinds: + writer = next((u for u in updates if _wrote(u, (*path, kind))), None) + if writer is not None: + writers.append(f"{kind} from {writer.source}") + if writers: + message += f"; {', '.join(writers)}" return message -def _locate( - model: type[BaseModel], location: Sequence[Any] -) -> tuple[list[str], list[str]]: - """Pydantic's error location as the names of a dotted path, and the keys - valid where it ends. A field name moves into its section; a kind moves - into its variant, whose tag field is not offered since the key names the - kind; a dict key or list index stays in the section; the class name - pydantic inserts under a `Section | scalar` field is dropped.""" - section: type[BaseModel] | None = model - kinds: dict[str, type[BaseModel]] | None = None - hidden: str | None = None - inside_collection = False - names = [] - for part in map(str, location): - if kinds is not None: - names.append(part) - section, kinds = kinds.get(part), None +def _override_without_effect( + update: Update, later: list[Update], merged: dict[str, Any] +) -> str | None: + """C11, the two ways an override loses its effect: a later override writes + at its path, above it, or below a value of its that was not a mapping (`{}` + never clears a mapping); or it wrote under a kind that does not run at its + kinds field, decided from the merged dict as `_to_tagged_form` decides it. + Never raises.""" + for other in reversed(later): + depth = min(len(update.path), len(other.path)) + if other.value == {} or update.path[:depth] != other.path[:depth]: continue - if inside_collection: - names.append(part) - inside_collection = False + if len(other.path) > len(update.path) and update.value == {}: continue - if section is not None and part == section.__name__: - continue - names.append(part) - if section is not None and part in section.model_fields: - field = section.model_fields[part] - hidden = _tag_of(field) - if hidden is not None: - kinds = _kinds_of(field) - else: - section, inside_collection = _section_of(field.annotation) - else: - section = None - candidates = list(section.model_fields) if section is not None else [] - return names, [c for c in candidates if c != hidden] - - -def _unknown_key_messages(model: type[BaseModel], error: ValidationError) -> list[str]: - """Pydantic's unknown-key errors as messages that name the full dotted - path and the closest valid key at that level.""" - messages = [] - for item in error.errors(): - # pydantic's error code for a key that matches no field (extra="forbid"). - # Every other code, e.g. a wrong type, is left for `load()` to re-raise. - if item["type"] != "extra_forbidden": - continue - *location, key = item["loc"] - names, candidates = _locate(model, location) - prefix = "".join(f"{name}." for name in names) - messages.append( - _unknown_key_message(f"{prefix}{key}", [f"{prefix}{c}" for c in candidates]) - ) - return messages + if other.selects is not None: + field, running = _dotted(other.selects), _as_typed(other.value) + return f"{update.source} is overridden: {field} runs {running}" + at, value = _dotted(other.path), _as_typed(other.value) + return f"{update.source} is overridden: {at} is {value}" + for field, kinds, key in update.under_kinds: + mapping: Any = merged + for step in field: + mapping = mapping.get(step) if isinstance(mapping, dict) else None + running = _selected_kind(mapping, kinds) if isinstance(mapping, dict) else None + if running is not None and running != key: + at = _dotted(field) + return f"{update.source}: {at}.{key} has no effect, {at} runs {running}" + return None From 91a155f8337e07fbc57b0cfd4fbf443c2990baf5 Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Tue, 22 Sep 2026 14:05:35 +0200 Subject: [PATCH 07/14] Export inf and nan as JSON constants instead of null (CORE-2 follow-up, #1556) pydantic's JSON mode writes inf and nan as null, so a config holding either stored a different value in the model metadata and could not be loaded back (a null is not a float). `ConfigSection` now sets `ser_json_inf_nan="constants"`, as `metadata._Record` already does: both exports keep them as floats, JSON writes Infinity and NaN, and JSON, YAML and TOML each read their own spelling back to the same config. From the PR #1733 review on e5339d4. The review's other comment, a schema behind a forward reference escaping the checks, is covered by rule F7 and its tests. --- .../mace-core/src/mace_core/config/base.py | 5 ++- .../mace-core/tests/test_mace_core_config.py | 32 +++++++++++++++++++ 2 files changed, 36 insertions(+), 1 deletion(-) diff --git a/packages/mace-core/src/mace_core/config/base.py b/packages/mace-core/src/mace_core/config/base.py index 69ff51cbd..514a95455 100644 --- a/packages/mace-core/src/mace_core/config/base.py +++ b/packages/mace-core/src/mace_core/config/base.py @@ -370,7 +370,10 @@ class ConfigSection(BaseModel): is held as given and exported under the variant's kind. """ - model_config = ConfigDict(extra="forbid") + # inf/nan are exported as floats (JSON constants Infinity, NaN) rather than + # pydantic's default null, which would store a different value in the model + # metadata; `metadata._Record` writes them the same way. + model_config = ConfigDict(extra="forbid", ser_json_inf_nan="constants") @classmethod def __pydantic_init_subclass__(cls, **kwargs: Any) -> None: diff --git a/packages/mace-core/tests/test_mace_core_config.py b/packages/mace-core/tests/test_mace_core_config.py index 35e913c10..47dd5a660 100644 --- a/packages/mace-core/tests/test_mace_core_config.py +++ b/packages/mace-core/tests/test_mace_core_config.py @@ -2,6 +2,7 @@ overrides without effect, and the resolved export's fixed point.""" import json +import math import re import subprocess import sys @@ -576,6 +577,37 @@ def test_fixed_point_holds_through_toml_when_nothing_is_none(tmp_path): assert_fixed_point(tmp_path, first, extension) +def test_inf_and_nan_survive_the_exports_in_every_format(tmp_path): + # pydantic's JSON mode writes them as null by default, which would put a + # different value into the model metadata. Each format spells them its own + # way (JSON constants, YAML .inf/.nan, TOML inf/nan); nan != nan, so the + # round trip is compared as JSON text. + config = DemoConfig.load( + write_config(tmp_path, ".yaml"), + ["--model.radial.cutoff", "inf", "--stage_two.energy_weight", "nan"], + ) + resolved = config.to_resolved_dict() + assert resolved["model"]["radial"]["cutoff"] == math.inf + assert math.isnan(resolved["stage_two"]["energy_weight"]) + assert math.isnan(config.to_user_dict()["stage_two"]["energy_weight"]) + text = json.dumps(resolved) + assert "null" not in text and "Infinity" in text and "NaN" in text + for extension, body in [ + (".json", text), + (".yaml", yaml.safe_dump(resolved)), + (".toml", "[model.radial]\ncutoff = inf\n[stage_two]\nenergy_weight = nan\n"), + ]: + path = tmp_path / f"special{extension}" + path.write_text(body, encoding="utf-8") + second = DemoConfig.load(path).to_resolved_dict() + assert second["model"]["radial"]["cutoff"] == math.inf, extension + assert math.isnan(second["stage_two"]["energy_weight"]), extension + assert ( + json.dumps(DemoConfig.load(tmp_path / "special.json").to_resolved_dict()) + == text + ) + + class LenientSection(BaseModel): cutoff: float = 5.0 From bf0f108188d110fff582e7af97734b4854cd28d0 Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Thu, 24 Sep 2026 12:41:05 +0200 Subject: [PATCH 08/14] Pin the revision-7 config contract in the tests before the base rewrite (CORE-2 follow-up, #1556) The config base is reduced to one file, one pydantic validation and two exports, and the command-line half moves to its own module. The tests now pin exactly that. `test_mace_core_config.py`: `load(config_file)` parses one TOML, YAML or JSON file and validates it; `from_dict(document)` is the same validation for a caller that edits the parsed dict first and does not write into it; everything the schema rejects, unknown keys included, is pydantic's `ValidationError` at its dotted location, and `ConfigError` is raised only for a file that cannot be read or parsed. `test_mace_core_config_kinds.py`: a kinds field is written in pydantic's tagged form and its errors land under the tag; both exports write the tag back, with the documented caveat that a variant built in code without its tag exports without `kind`. The definition-time checks keep the set, alias, excluded and computed-field rules, the lenient section rule (also behind a forward reference) and the `extra="forbid"` guard. `test_mace_core_config_cli.py` is new and pins the command-line half: `parse_overrides` reads `--a.b value` and `--a.b=value` tokens into a mapping of dotted paths (`null`, `[` and `{` values are JSON, anything else a string for pydantic to coerce; a bad token, a missing value or bad JSON is a `ConfigError`), and `apply_overrides` writes such a mapping into a copy of the parsed file set-at-path, creating mappings on the way, replacing a parent that is not a mapping, merging a mapping value into a mapping and replacing otherwise, without writing into either argument, so a YAML anchor's two keys stop sharing one object and an anchor that contains itself is a `ConfigError`. Composed with `read_config_file` and `from_dict`, an override beats the file, which beats the defaults, and an unknown key in an override is pydantic's error at the path the override named. Gone with the contract: several files, merge and warning rules, `ConfigWarning`, the kind-as-key wire form, bare kind names, kind switching, the union-shape schema rules and the unknown-key neighbour messages. Tests revision 7 does not contradict are kept, some renamed or reworded. These tests fail on purpose against the module on the branch; the next commit replaces it. --- .../mace-core/tests/test_mace_core_config.py | 457 ++------ .../tests/test_mace_core_config_cli.py | 294 ++++++ .../tests/test_mace_core_config_kinds.py | 989 +++--------------- 3 files changed, 525 insertions(+), 1215 deletions(-) create mode 100644 packages/mace-core/tests/test_mace_core_config_cli.py diff --git a/packages/mace-core/tests/test_mace_core_config.py b/packages/mace-core/tests/test_mace_core_config.py index 47dd5a660..c9414d477 100644 --- a/packages/mace-core/tests/test_mace_core_config.py +++ b/packages/mace-core/tests/test_mace_core_config.py @@ -1,12 +1,13 @@ -"""`ReforgeBaseConfig`: file formats, precedence, dotted overrides, unknown keys, -overrides without effect, and the resolved export's fixed point.""" +"""`ReforgeBaseConfig`: file formats, `from_dict`, unknown keys, and the two +exports' fixed point.""" +import inspect import json import math import re import subprocess import sys -import warnings +from collections.abc import Set as AbstractSet from typing import Annotated, Any import pytest @@ -14,8 +15,8 @@ from mace_core.config import ( ConfigError, ConfigSection, - ConfigWarning, ReforgeBaseConfig, + read_config_file, ) from pydantic import BaseModel, ConfigDict, Field, ValidationError, computed_field @@ -94,6 +95,11 @@ def write_config(tmp_path, extension, values=FILE_VALUES, name="config"): return path +def error_locations(excinfo): + """The dotted location of every error a `ValidationError` carries.""" + return [".".join(map(str, error["loc"])) for error in excinfo.value.errors()] + + # --------------------------------------------------------------------------- # File loading @@ -107,6 +113,10 @@ def test_same_config_loads_identically_from_every_format(tmp_path, extension): assert config.model.radial.num_bessel == 8 +def test_read_config_file_returns_the_parsed_dict(tmp_path): + assert read_config_file(write_config(tmp_path, ".toml")) == FILE_VALUES + + def test_extension_is_matched_in_any_case(tmp_path): path = tmp_path / "CONFIG.YAML" path.write_text("seed: 5\n", encoding="utf-8") @@ -132,10 +142,10 @@ def test_comment_only_file_is_all_defaults(tmp_path): assert DemoConfig.load(path) == DemoConfig() -def test_file_must_be_a_table_at_the_top(tmp_path): +def test_file_must_be_a_mapping_at_the_top(tmp_path): path = tmp_path / "list.json" path.write_text("[1, 2]", encoding="utf-8") - with pytest.raises(ConfigError, match="table of keys at the top level"): + with pytest.raises(ConfigError, match="mapping of keys to values at the top"): DemoConfig.load(path) @@ -162,16 +172,22 @@ def test_malformed_file_is_a_config_error(tmp_path, extension, text): DemoConfig.load(path) -# --------------------------------------------------------------------------- -# Precedence: defaults < files in order < CLI. The legacy behaviour this pins is -# tests/unit/test_arg_parser.py::test_cli_flag_overrides_yaml_config. +def test_a_yaml_anchor_that_contains_itself_is_a_config_error(tmp_path): + # Under a section it would be pydantic's error; under a free dict it would + # load and then fail to export, so the file is refused up front. + class Free(ReforgeBaseConfig): + extra: dict[str, Any] = Field(default_factory=dict) + path = tmp_path / "loop.yaml" + path.write_text("extra: &loop {b: *loop}\n", encoding="utf-8") + with pytest.raises(ConfigError, match=r"loop\.yaml contains a value that refers"): + Free.load(path) + path.write_text("extra: {a: &shared {x: 1}, b: *shared}\n", encoding="utf-8") + assert Free.load(path).extra == {"a": {"x": 1}, "b": {"x": 1}} # sharing is fine -def test_no_inputs_gives_the_defaults(): - config = DemoConfig.load() - assert config == DemoConfig() - assert config.model.num_interactions == 2 - assert DemoConfig.load(None) == DemoConfig() # an optional path, unset + +# --------------------------------------------------------------------------- +# Precedence: defaults < the file. Nothing else feeds a config. def test_file_overrides_defaults(tmp_path): @@ -180,316 +196,86 @@ def test_file_overrides_defaults(tmp_path): assert config.default_dtype == "float64" # untouched default -def test_cli_overrides_file_which_overrides_defaults(tmp_path): - config = DemoConfig.load( - write_config(tmp_path, ".toml"), ["--model.num_interactions", "3"] - ) - assert config.model.num_interactions == 3 # CLI beats the file's 4 - assert config.model.radial.cutoff == 4.5 # the file's other values survive - assert config.seed == 7 - assert config.model.radial.num_bessel == 8 # defaults fill the rest - assert config.default_dtype == "float64" - - -def test_files_apply_in_order_before_the_overrides(tmp_path): - first = write_config(tmp_path, ".yaml", name="defaults") - second = write_config( - tmp_path, ".toml", {"seed": 8, "model": {"radial": {"num_bessel": 6}}}, "site" - ) - config = DemoConfig.load([first, second], ["--model.num_interactions", "3"]) - assert config.seed == 8 # the second file beats the first - assert config.name == "water" # the first file's other values survive - assert config.model.radial == RadialSection(num_bessel=6, cutoff=4.5) - assert config.model.num_interactions == 3 # the CLI beats both - assert DemoConfig.load([]) == DemoConfig() - - -# --------------------------------------------------------------------------- -# Dotted CLI overrides - - -def test_dotted_override_reaches_a_two_level_nested_field(): - config = DemoConfig.load(cli_overrides=["--model.radial.cutoff", "6.0"]) - assert config.model.radial.cutoff == 6.0 - assert config.model.radial.num_bessel == 8 - - -def test_dotted_override_reaches_a_section_left_at_its_defaults(): - config = DemoConfig.load(cli_overrides=["--stage_two.start_epoch", "50"]) - assert config.stage_two == StageTwoSection(start_epoch=50) - assert DemoConfig.load().stage_two == StageTwoSection() - - -def test_override_forms_and_types(): - config = DemoConfig.load( - cli_overrides=[ - "--seed=9", - "--data.train_file", - "null", - "--data.heads", - '["a", "b"]', - ] - ) - assert config.seed == 9 +def test_values_are_handed_to_pydantic_as_given(tmp_path): + values = {"seed": "9", "data": {"train_file": None, "heads": ["a", "b"]}} + config = DemoConfig.load(write_config(tmp_path, ".json", values)) + assert config.seed == 9 # pydantic's lax coercion, not the loader's assert config.data.train_file is None assert config.data.heads == ["a", "b"] def test_value_of_the_wrong_type_is_a_validation_error(tmp_path): - with pytest.raises(ValidationError, match="seed"): - DemoConfig.load(cli_overrides=["--seed", "seven"]) with pytest.raises(ValidationError, match="seed"): DemoConfig.load(write_config(tmp_path, ".yaml", {"seed": "seven"})) -def test_override_missing_its_value_is_a_config_error(): - with pytest.raises(ConfigError, match="override --seed is missing its value"): - DemoConfig.load(cli_overrides=["--seed"]) - - -def test_override_that_is_not_valid_json_is_a_config_error(): - with pytest.raises(ConfigError, match="override --model is not valid JSON"): - DemoConfig.load(cli_overrides=["--model", "{oops"]) - - -def test_value_starting_with_dashes_works_in_both_forms(): - assert DemoConfig.load(cli_overrides=["--name=--odd"]).name == "--odd" - assert DemoConfig.load(cli_overrides=["--name", "--odd"]).name == "--odd" - - -@pytest.mark.parametrize("token", ["--", "--=5"]) -def test_bare_dashes_are_an_unknown_option_not_a_key(token): - with pytest.raises(ConfigError, match=rf"unknown config option '{token}'"): - DemoConfig.load(cli_overrides=[token, "--seed", "5"]) - - -def test_overrides_given_as_one_string_are_refused(): - with pytest.raises(TypeError, match="cli_overrides is a string"): - DemoConfig.load(cli_overrides="--seed 5") - - -def test_dict_valued_field_takes_json_and_dotted_paths_into_its_entries(): - class Sources(ReforgeBaseConfig): - by_name: dict[str, RadialSection] = Field(default_factory=dict) - - config = Sources.load(cli_overrides=["--by_name", '{"pbe": {"cutoff": 4.0}}']) - assert config.by_name == {"pbe": RadialSection(cutoff=4.0)} - dotted = Sources.load(cli_overrides=["--by_name.pbe.cutoff", "4.0"]) - assert dotted.by_name == {"pbe": RadialSection(cutoff=4.0)} - # Inside an entry, the neighbour is still found: the key passes through. - with pytest.raises( - ConfigError, - match=r"'by_name\.pbe\.cutof'; did you mean 'by_name\.pbe\.cutoff'\?", - ): - Sources.load(cli_overrides=["--by_name", '{"pbe": {"cutof": 4.0}}']) - - -def test_collections_behind_none_or_annotated_keep_their_dotted_paths(tmp_path): - class Collections(ReforgeBaseConfig): - counts: dict[str, int] | None = None - documented: dict[str, Annotated[RadialSection, Field(description="d")]] = Field( - default_factory=dict - ) - pair: tuple[int, RadialSection] | None = None - - assert Collections.load(cli_overrides=["--counts.x", "1"]).counts == {"x": 1} - with pytest.raises( - ConfigError, - match=r"'documented\.a\.cutof'; did you mean 'documented\.a\.cutoff'", - ): - Collections.load(cli_overrides=["--documented", '{"a": {"cutof": 4.0}}']) - path = write_config(tmp_path, ".json", {"pair": [1, {"cutof": 4.0}]}) - with pytest.raises( - ConfigError, match=r"'pair\.1\.cutof'; did you mean 'pair\.1\.cutoff'" - ): - Collections.load(path) - - -def test_dict_override_merges_entries_but_list_override_replaces(tmp_path): - class Sources(ReforgeBaseConfig): - by_name: dict[str, RadialSection] = Field(default_factory=dict) - heads: list[str] = Field(default_factory=list) - - path = write_config( - tmp_path, ".yaml", {"by_name": {"pbe": {"cutoff": 4.0}}, "heads": ["a", "b"]} - ) - config = Sources.load( - path, ["--by_name", '{"r2scan": {"cutoff": 6.0}}', "--heads", '["c"]'] - ) - assert set(config.by_name) == {"pbe", "r2scan"} - assert config.heads == ["c"] - - -def test_overrides_apply_in_order_on_top_of_the_file(tmp_path): - # A dotted value merges into what the file set in the section; a dotted - # value followed by the whole section keeps both. - config = DemoConfig.load( - write_config(tmp_path, ".yaml", {"stage_two": {"energy_weight": 5.0}}), - cli_overrides=[ - "--stage_two.start_epoch", - "5", - "--model.radial.cutoff", - "4", - "--model", - '{"num_interactions": 3}', - ], - ) - assert config.stage_two == StageTwoSection(start_epoch=5, energy_weight=5.0) - assert (config.model.num_interactions, config.model.radial.cutoff) == (3, 4.0) - - -def test_a_yaml_anchor_does_not_share_an_override(tmp_path): - class Two(ReforgeBaseConfig): - a: dict[str, Any] = Field(default_factory=dict) - b: dict[str, Any] = Field(default_factory=dict) - - path = tmp_path / "anchors.yaml" - path.write_text("a: &empty {}\nb: *empty\n", encoding="utf-8") - config = Two.load(path, ["--a.x", "1"]) - assert (config.a, config.b) == ({"x": "1"}, {}) - - # --------------------------------------------------------------------------- -# An override that had no effect on the config that runs is a warning; a file -# never warns. - - -class Extras(ReforgeBaseConfig): - seed: int = 1 - extra: dict[str, Any] = Field(default_factory=dict) - - -WITHOUT_EFFECT = { - "a later override at the same path": ( - ["--seed", "1", "--seed", "2"], - ["--seed 1 is overridden: seed is 2"], - ), - "a later override above it": ( - ["--extra.a.b", "2", "--extra.a", "5"], - ["--extra.a.b 2 is overridden: extra.a is 5"], - ), - "a later json override above it": ( - ["--extra.a.b", "2", "--extra.a", "[1]"], - ["--extra.a.b 2 is overridden: extra.a is [1]"], - ), - "part of a json override replaced": ( - ["--extra", '{"a": {"b": 1}}', "--extra.a.b", "2"], - ['--extra {"a": {"b": 1}} is overridden: extra.a.b is 2'], - ), - "a later override below a scalar of it": ( - ["--extra.a", "5", "--extra.a.b", "2"], - ["--extra.a 5 is overridden: extra.a.b is 2"], - ), - "the same override twice: the first is overridden": ( - ["--seed", "2", "--seed", "2"], - ["--seed 2 is overridden: seed is 2"], - ), - "the same json override twice: the first is overridden": ( - ["--extra", '{"a": 1}', "--extra", '{"a": 1}'], - ['--extra {"a": 1} is overridden: extra.a is 1'], - ), - "the last of three at one path is named, once": ( - ["--seed", "1", "--seed", "1", "--seed", "3"], - ["--seed 1 is overridden: seed is 3"], - ), - "a plain key named kind is a key, not a selection": ( - ["--extra.kind", "a", "--extra.kind", "b"], - ["--extra.kind a is overridden: extra.kind is b"], - ), - "an empty mapping above a later key is a merge": ( - ["--extra.a.b", "2", "--extra.a", "{}"], - [], - ), - "an empty mapping below a later key is a merge": ( - ["--extra.a", "{}", "--extra.a.b", "2"], - [], - ), -} +# `from_dict` is `load` for a caller that edits the parsed file first (a +# command line writing its flags); it validates the same way. -@pytest.mark.parametrize("row", WITHOUT_EFFECT, ids=WITHOUT_EFFECT) -def test_an_override_without_effect_warns(row): - tokens, expected = WITHOUT_EFFECT[row] - with warnings.catch_warnings(record=True) as caught: - warnings.simplefilter("always") - Extras.load(cli_overrides=tokens) - assert [str(w.message) for w in caught] == expected - assert all(issubclass(w.category, ConfigWarning) for w in caught) +def test_from_dict_validates_a_parsed_document(tmp_path): + document = read_config_file(write_config(tmp_path, ".toml")) + document["seed"] = 9 + document.setdefault("stage_two", {})["start_epoch"] = 50 + config = DemoConfig.from_dict(document) + assert config.seed == 9 + assert config.model.num_interactions == 4 # the file's values survive + assert config.stage_two == StageTwoSection(start_epoch=50) + assert DemoConfig.load(write_config(tmp_path, ".toml")) == DemoConfig.from_dict( + read_config_file(write_config(tmp_path, ".toml")) + ) -def test_a_file_value_the_cli_replaces_does_not_warn(tmp_path): - with warnings.catch_warnings(): - warnings.simplefilter("error", ConfigWarning) - config = DemoConfig.load(write_config(tmp_path, ".yaml"), ["--seed", "9"]) - assert config.seed == 9 +def test_from_dict_does_not_write_into_the_document(): + document = {"model": {"radial": {"cutoff": 6.0}}} + DemoConfig.from_dict(document) + assert document == {"model": {"radial": {"cutoff": 6.0}}} # --------------------------------------------------------------------------- -# Unknown keys name the key and its nearest neighbour, in files and on the CLI. +# Unknown keys are pydantic's error, at their dotted location; every error of +# one load is reported together. def test_unknown_top_level_key_in_file(tmp_path): path = tmp_path / "typo.yaml" path.write_text("sead: 1\n", encoding="utf-8") - with pytest.raises(ConfigError, match=r"'sead'; did you mean 'seed'\?"): + with pytest.raises(ValidationError) as excinfo: DemoConfig.load(path) + assert error_locations(excinfo) == ["sead"] + assert excinfo.value.errors()[0]["type"] == "extra_forbidden" def test_unknown_nested_key_in_file_names_the_dotted_path(tmp_path): path = tmp_path / "typo.json" path.write_text(json.dumps({"model": {"radial": {"cutof": 4.0}}}), encoding="utf-8") - with pytest.raises( - ConfigError, - match=r"'model\.radial\.cutof'; did you mean 'model\.radial\.cutoff'\?", - ): + with pytest.raises(ValidationError) as excinfo: DemoConfig.load(path) + assert error_locations(excinfo) == ["model.radial.cutof"] + assert "model.radial.cutof" in str(excinfo.value) -def test_every_unknown_key_is_reported_at_once(tmp_path): +def test_every_error_is_reported_at_once(tmp_path): path = tmp_path / "typos.yaml" - path.write_text("sead: 1\nmodel:\n num_interaction: 3\n", encoding="utf-8") - with pytest.raises(ConfigError) as excinfo: - DemoConfig.load(path) - assert "'sead'" in str(excinfo.value) - assert "'model.num_interaction'" in str(excinfo.value) - - -def test_unknown_keys_are_reported_under_their_file_or_override(tmp_path): - first = write_config(tmp_path, ".yaml", {"sead": 1}, "first") - second = write_config( - tmp_path, ".json", {"model": {"num_interaction": 3}}, "second" + path.write_text( + "sead: 1\nseed: seven\nmodel:\n num_interaction: 3\n", encoding="utf-8" ) - with pytest.raises(ConfigError) as excinfo: - DemoConfig.load([first, second], ["--nmae", "x"]) - assert str(excinfo.value).splitlines() == [ - f"{first}: unknown config key 'sead'; did you mean 'seed'?", - f"{second}: unknown config key 'model.num_interaction'; " - "did you mean 'model.num_interactions'?", - "--nmae x: unknown config key 'nmae'; did you mean 'name'?", - ] - - -def test_unknown_key_without_a_close_neighbour_still_names_it(tmp_path): - path = tmp_path / "far.yaml" - path.write_text("zzzzzz: 1\n", encoding="utf-8") - with pytest.raises(ConfigError, match=r"unknown config key 'zzzzzz'$"): + with pytest.raises(ValidationError) as excinfo: DemoConfig.load(path) + assert set(error_locations(excinfo)) == {"sead", "seed", "model.num_interaction"} -def test_unknown_dotted_override_names_the_neighbour(): - with pytest.raises( - ConfigError, - match=r"'model\.num_interaction'; did you mean 'model\.num_interactions'\?", - ): - DemoConfig.load(cli_overrides=["--model.num_interaction", "3"]) - - -def test_unknown_key_inside_a_nested_section(): - with pytest.raises( - ConfigError, - match=r"'stage_two\.start'; did you mean 'stage_two\.start_epoch'\?", - ): - DemoConfig.load(cli_overrides=["--stage_two.start", "50"]) +def test_unknown_key_in_a_document_is_reported_at_its_path(): + documents = { + "nmae": {"nmae": 1}, + "model.num_interaction": {"model": {"num_interaction": 1}}, + "stage_two.start": {"stage_two": {"start": 1}}, + } + for dotted_path, document in documents.items(): + with pytest.raises(ValidationError) as excinfo: + DemoConfig.from_dict(document) + assert error_locations(excinfo) == [dotted_path] def test_every_bad_list_item_is_reported(tmp_path): @@ -497,27 +283,9 @@ class Layers(ReforgeBaseConfig): layers: list[RadialSection] = Field(default_factory=list) path = write_config(tmp_path, ".json", {"layers": [{"cutof": 1}, {"nb": 2}]}) - with pytest.raises(ConfigError) as excinfo: + with pytest.raises(ValidationError) as excinfo: Layers.load(path) - assert re.findall(r"unknown config key '([^']*)'", str(excinfo.value)) == [ - "layers.0.cutof", - "layers.1.nb", - ] - - -def test_help_flag_is_an_error_not_an_exit(): - with pytest.raises(ConfigError, match=r"unknown config option '-h'"): - DemoConfig.load(cli_overrides=["-h"]) - - -def test_abbreviated_option_is_unknown_not_expanded(): - with pytest.raises(ConfigError, match=r"key 'se'; did you mean 'seed'"): - DemoConfig.load(cli_overrides=["--se", "3"]) - - -def test_empty_inline_value_does_not_hide_the_next_option(): - with pytest.raises(ConfigError, match=r"'sead'; did you mean 'seed'\?"): - DemoConfig.load(cli_overrides=["--name=", "--sead", "5"]) + assert error_locations(excinfo) == ["layers.0.cutof", "layers.1.nb"] def test_direct_construction_rejects_unknown_keys_too(): @@ -558,10 +326,8 @@ def assert_fixed_point(tmp_path, first, extension): @pytest.mark.parametrize("extension", [".yaml", ".json"]) def test_file_to_resolved_to_file_to_resolved_is_a_fixed_point(tmp_path, extension): - first = DemoConfig.load( - write_config(tmp_path, ".toml"), - ["--model.num_interactions", "3", "--data.train_file", "null"], - ).to_resolved_dict() + values = {**FILE_VALUES, "data": {"train_file": None, "heads": ["pbe"]}} + first = DemoConfig.load(write_config(tmp_path, ".yaml", values)).to_resolved_dict() assert first["data"]["train_file"] is None # a None is part of what has to survive assert_fixed_point(tmp_path, first, extension) @@ -569,9 +335,8 @@ def test_file_to_resolved_to_file_to_resolved_is_a_fixed_point(tmp_path, extensi def test_fixed_point_holds_through_toml_when_nothing_is_none(tmp_path): # TOML has no null, so the file sets the optional file name; the resolved # dict then goes through all three formats. - first = DemoConfig.load( - write_config(tmp_path, ".yaml"), ["--stage_two.start_epoch", "50"] - ).to_resolved_dict() + values = {**FILE_VALUES, "stage_two": {"start_epoch": 50}} + first = DemoConfig.load(write_config(tmp_path, ".yaml", values)).to_resolved_dict() assert "null" not in json.dumps(first) for extension in (".toml", ".yaml", ".json"): assert_fixed_point(tmp_path, first, extension) @@ -582,10 +347,12 @@ def test_inf_and_nan_survive_the_exports_in_every_format(tmp_path): # different value into the model metadata. Each format spells them its own # way (JSON constants, YAML .inf/.nan, TOML inf/nan); nan != nan, so the # round trip is compared as JSON text. - config = DemoConfig.load( - write_config(tmp_path, ".yaml"), - ["--model.radial.cutoff", "inf", "--stage_two.energy_weight", "nan"], - ) + values = { + **FILE_VALUES, # the file sets the optional file name, so no null + "model": {"radial": {"cutoff": math.inf}}, + "stage_two": {"energy_weight": math.nan}, + } + config = DemoConfig.load(write_config(tmp_path, ".yaml", values)) resolved = config.to_resolved_dict() assert resolved["model"]["radial"]["cutoff"] == math.inf assert math.isnan(resolved["stage_two"]["energy_weight"]) @@ -615,10 +382,10 @@ class LenientSection(BaseModel): def test_field_shapes_the_contract_cannot_keep_are_rejected_at_class_definition(): # Each shape would break a guarantee: set order varies with the hash # seed; aliases, excluded and computed fields do not validate back; a - # union of sections would let a value pick its section; a lenient section - # would swallow typos; a section is never optional. + # lenient section would swallow typos. shapes = { r"tags is typed as a set.*Use a list": ("tags", list[set[str]]), + r"names is typed as a set.*Use a list": ("names", AbstractSet[str]), r"num has an alias": ("num", Annotated[int, Field(alias="n")]), r"vnum has an alias": ("vnum", Annotated[int, Field(validation_alias="n")]), r"snum has an alias": ("snum", Annotated[int, Field(serialization_alias="n")]), @@ -626,13 +393,10 @@ def test_field_shapes_the_contract_cannot_keep_are_rejected_at_class_definition( "hidden", Annotated[int, Field(exclude=True)], ), - r"either is a union of sections": ("either", RadialSection | StageTwoSection), r"radial holds LenientSection, which is not a ConfigSection": ( "radial", LenientSection | None, ), - r"radial mixes its kinds with int": ("radial", RadialSection | int), - r"stage admits None": ("stage", StageTwoSection | None), } for message, (name, annotation) in shapes.items(): with pytest.raises(TypeError, match=message): @@ -685,44 +449,39 @@ class LeakingConfig(ReforgeBaseConfig): leaking: Leaking = Field(default_factory=Leaking) -def test_a_forward_reference_is_checked_and_walked_once_it_resolves(tmp_path): - assert not Forward.__pydantic_complete__ - assert ForwardConfig.load().forward.later == Later() - config = ForwardConfig.load(cli_overrides=["--forward.later.x", "2"]) +def test_a_forward_reference_is_resolved_and_checked_on_load(tmp_path): + empty = write_config(tmp_path, ".yaml", {}) + assert ForwardConfig.load(empty).forward.later == Later() + config = ForwardConfig.from_dict({"forward": {"later": {"x": 2}}}) assert config.forward.later == Later(x=2) path = tmp_path / "typo.json" path.write_text(json.dumps({"forward": {"later": {"x": 2, "typo": 1}}})) - with pytest.raises(ConfigError, match=r"unknown config key 'forward.later.typo'"): + with pytest.raises(ValidationError) as excinfo: ForwardConfig.load(path) + assert error_locations(excinfo) == ["forward.later.typo"] -def test_a_lenient_section_behind_a_forward_reference_is_rejected(): +def test_a_lenient_section_behind_a_forward_reference_is_rejected(tmp_path): with pytest.raises( TypeError, match=r"Leaking.plain holds PlainLater, which is not a ConfigSection" ): - LeakingConfig.load() + LeakingConfig.load(write_config(tmp_path, ".yaml", {})) def test_user_dict_holds_only_what_was_set(tmp_path): - config = DemoConfig.load( - write_config(tmp_path, ".json"), ["--model.num_interactions", "3"] - ) - assert config.to_user_dict() == { - "name": "water", - "seed": 7, - "model": {"num_interactions": 3, "radial": {"cutoff": 4.5}}, - "data": {"train_file": "train.xyz", "heads": ["pbe", "r2scan"]}, - } + config = DemoConfig.load(write_config(tmp_path, ".json")) + assert config.to_user_dict() == FILE_VALUES + assert config.to_user_dict() is not FILE_VALUES # --------------------------------------------------------------------------- -# Nothing but the files and the CLI feed a config. +# Nothing but the file feeds a config. -def test_environment_variables_are_ignored(monkeypatch): +def test_environment_variables_are_ignored(tmp_path, monkeypatch): monkeypatch.setenv("NAME", "from-the-environment") monkeypatch.setenv("SEED", "99") - config = DemoConfig.load() + config = DemoConfig.load(write_config(tmp_path, ".yaml", {})) assert config.name == "mace" assert config.seed == 123 @@ -733,8 +492,9 @@ def test_root_keys_are_validated_like_any_section(tmp_path, key): # options are special at the top level: unknown is unknown. path = tmp_path / "root.json" path.write_text(json.dumps({key: 1}), encoding="utf-8") - with pytest.raises(ConfigError, match=f"unknown config key '{key}'"): + with pytest.raises(ValidationError) as excinfo: DemoConfig.load(path) + assert error_locations(excinfo) == [key] def test_config_module_imports_neither_torch_nor_jax(): @@ -745,3 +505,12 @@ def test_config_module_imports_neither_torch_nor_jax(): "assert not leaked, leaked" ) subprocess.run([sys.executable, "-c", code], check=True) + + +def test_the_base_does_not_import_the_command_line_module(): + # The package init re-exports both, so `sys.modules` cannot tell; the + # source can: `cli` imports `ConfigError` from `base`, never the reverse. + import mace_core.config.base as base_module + + source = inspect.getsource(base_module) + assert not re.search(r"^\s*(from|import) .*\bcli\b", source, re.MULTILINE) diff --git a/packages/mace-core/tests/test_mace_core_config_cli.py b/packages/mace-core/tests/test_mace_core_config_cli.py new file mode 100644 index 000000000..ee7b1d112 --- /dev/null +++ b/packages/mace-core/tests/test_mace_core_config_cli.py @@ -0,0 +1,294 @@ +"""`mace_core.config.cli`: the `--a.b value` grammar to a mapping of dotted +paths, and that mapping written into a parsed file before one validation. The +config base knows nothing of either; a command line composes them with +`read_config_file` and `from_dict`.""" + +import json +from typing import Annotated, Any, Literal + +import pytest +from mace_core.config import ( + ConfigError, + ConfigSection, + ReforgeBaseConfig, + apply_overrides, + parse_overrides, + read_config_file, +) +from pydantic import Field, ValidationError + +#: A warning the test did not ask for is a failure. +pytestmark = pytest.mark.filterwarnings("error") + +# --------------------------------------------------------------------------- +# The demo schema: two levels of nesting, a list, an optional, a free dict, a +# kinds field. + + +class RadialSection(ConfigSection): + num_bessel: int = 8 + cutoff: float = 5.0 + + +class ModelSection(ConfigSection): + num_interactions: int = 2 + radial: RadialSection = RadialSection() + + +class DataSection(ConfigSection): + train_file: str | None = None + heads: list[str] = Field(default_factory=lambda: ["default"]) + + +class StageTwoSection(ConfigSection): + start_epoch: int = 100 + energy_weight: float = 1000.0 + + +class Weighted(ConfigSection): + kind: Literal["weighted"] = "weighted" + stress_weight: float = 0.0 + + +class Huber(ConfigSection): + kind: Literal["huber"] = "huber" + delta: float = 0.01 + + +class DemoConfig(ReforgeBaseConfig): + name: str = "mace" + seed: int = 123 + model: ModelSection = ModelSection() + data: DataSection = DataSection() + stage_two: StageTwoSection = StageTwoSection() + loss: Annotated[Weighted | Huber, Field(discriminator="kind")] = Weighted() + extra: dict[str, Any] = Field(default_factory=dict) + + +FILE_VALUES = { + "name": "water", + "seed": 7, + "model": {"num_interactions": 4, "radial": {"cutoff": 4.5}}, + "data": {"train_file": "train.xyz", "heads": ["pbe", "r2scan"]}, +} + +HEADS_JSON = '["a", "b"]' + + +def write_config(tmp_path, values=FILE_VALUES): + path = tmp_path / "config.json" + path.write_text(json.dumps(values), encoding="utf-8") + return path + + +def load(tmp_path, argv, values=FILE_VALUES, root=DemoConfig): + """What a command line does with its config file and the tokens after it.""" + document = read_config_file(write_config(tmp_path, values)) + return root.from_dict(apply_overrides(document, parse_overrides(argv))) + + +def error_locations(excinfo): + return [".".join(map(str, error["loc"])) for error in excinfo.value.errors()] + + +# --------------------------------------------------------------------------- +# parse_overrides: tokens to a mapping of dotted path to value. + + +def test_both_forms_and_the_value_types(): + argv = ["--seed=9", "--data.train_file", "null", "--data.heads", HEADS_JSON] + assert parse_overrides(argv) == { + "seed": "9", # a string: pydantic coerces it, the grammar does not + "data.train_file": None, + "data.heads": ["a", "b"], + } + assert parse_overrides(["--model", json.dumps({"num_interactions": 3})]) == { + "model": {"num_interactions": 3} + } + assert parse_overrides([]) == {} + + +def test_paths_keep_their_order_and_a_repeated_path_moves_to_where_it_is_last(): + overrides = parse_overrides(["--b", "1", "--a", "2", "--b", "3"]) + assert list(overrides.items()) == [("a", "2"), ("b", "3")] + # So a later write under a parent the same command line replaced survives. + argv = ["--m.r.c", "4", "--m.r", "null", "--m.r.c", "5"] + assert apply_overrides({}, parse_overrides(argv)) == {"m": {"r": {"c": "5"}}} + + +def test_an_empty_inline_value_does_not_hide_the_next_option(): + assert parse_overrides(["--name=", "--seed", "5"]) == {"name": "", "seed": "5"} + + +@pytest.mark.parametrize("token", ["--a..b", "--a.", "--.a"]) +def test_an_empty_key_in_a_path_is_a_config_error(token): + with pytest.raises(ConfigError, match=rf"override {token} has an empty key"): + parse_overrides([token, "1"]) + + +def test_a_value_starting_with_dashes_works_in_both_forms(): + assert parse_overrides(["--name=--odd"]) == {"name": "--odd"} + assert parse_overrides(["--name", "--odd"]) == {"name": "--odd"} + + +def test_a_path_missing_its_value_is_a_config_error(): + with pytest.raises(ConfigError, match="override --seed is missing its value"): + parse_overrides(["--seed"]) + + +def test_a_value_that_is_not_valid_json_is_a_config_error(): + with pytest.raises(ConfigError, match="override --model is not valid JSON"): + parse_overrides(["--model", "{oops"]) + + +@pytest.mark.parametrize("token", ["--", "--=5", "seed=5", "-s"]) +def test_a_token_that_is_not_a_dotted_option_is_a_config_error(token): + with pytest.raises(ConfigError, match=rf"unknown config option '{token}'"): + parse_overrides([token, "--seed", "5"]) + + +def test_tokens_given_as_one_string_are_refused(): + with pytest.raises(TypeError, match="tokens is a string"): + parse_overrides("--seed 5") + + +# --------------------------------------------------------------------------- +# apply_overrides: a dotted path is written into a copy of the parsed file. + + +def test_a_top_level_key_is_written(): + assert apply_overrides({}, {"seed": 9}) == {"seed": 9} + assert apply_overrides({"name": "x"}, {"seed": 9}) == {"name": "x", "seed": 9} + + +def test_a_nested_path_writes_into_the_section_the_file_set(): + document = {"model": {"num_interactions": 4, "radial": {"cutoff": 4.5}}} + assert apply_overrides(document, {"model.radial.cutoff": 6.0}) == { + "model": {"num_interactions": 4, "radial": {"cutoff": 6.0}} + } + + +def test_the_sections_a_path_passes_through_are_created(): + assert apply_overrides({}, {"model.radial.cutoff": 6.0, "stage_two.x": 1}) == { + "model": {"radial": {"cutoff": 6.0}}, + "stage_two": {"x": 1}, + } + + +@pytest.mark.parametrize("parent", [5, [1, 2], None], ids=["scalar", "list", "null"]) +def test_a_parent_that_is_not_a_mapping_is_replaced(parent): + assert apply_overrides({"a": parent}, {"a.b": 2}) == {"a": {"b": 2}} + + +def test_a_mapping_value_merges_into_a_mapping_but_a_list_replaces(): + document = {"by_name": {"pbe": {"cutoff": 4.0}}, "heads": ["a", "b"]} + merged = apply_overrides( + document, {"by_name": {"r2scan": {"cutoff": 6.0}}, "heads": ["c"]} + ) + assert merged == { + "by_name": {"pbe": {"cutoff": 4.0}, "r2scan": {"cutoff": 6.0}}, + "heads": ["c"], + } + # Deeper too, and a key of the JSON value is one key even with a dot in it. + assert apply_overrides( + {"a": {"b": {"c": 1, "d": 2}}}, {"a": {"b": {"c": 3}, "x.y": 4}} + ) == {"a": {"b": {"c": 3, "d": 2}, "x.y": 4}} + # A mapping over a scalar, or a scalar over a mapping, replaces. + assert apply_overrides({"a": 1}, {"a": {"b": 2}}) == {"a": {"b": 2}} + assert apply_overrides({"a": {"b": 2}}, {"a": 1}) == {"a": 1} + + +def test_paths_apply_in_order_on_top_of_the_file(): + # A dotted value followed by the whole section keeps both. + document = {"stage_two": {"energy_weight": 5.0}} + argv = ["--stage_two.start_epoch", "5", "--model.radial.cutoff", "4"] + argv += ["--model", json.dumps({"num_interactions": 3})] + assert apply_overrides(document, parse_overrides(argv)) == { + "stage_two": {"energy_weight": 5.0, "start_epoch": "5"}, + "model": {"radial": {"cutoff": "4"}, "num_interactions": 3}, + } + + +def test_neither_the_document_nor_the_overrides_are_written_into(): + document = {"extra": {"a": {"x": {}}}} + overrides = {"extra.a": {"x": {"y": 1}}, "extra.a.x.z": [1]} + copy = apply_overrides(document, overrides) + assert copy == {"extra": {"a": {"x": {"y": 1, "z": [1]}}}} + assert document == {"extra": {"a": {"x": {}}}} + assert overrides == {"extra.a": {"x": {"y": 1}}, "extra.a.x.z": [1]} + copy["extra"]["a"]["x"]["z"].append(2) # the value was copied too + assert overrides["extra.a.x.z"] == [1] + + +def test_a_yaml_anchor_does_not_share_an_override(tmp_path): + path = tmp_path / "anchors.yaml" + path.write_text("a: &empty {}\nb: *empty\n", encoding="utf-8") + document = read_config_file(path) + assert document["a"] is document["b"] # what the parser hands over + assert apply_overrides(document, {"a.x": 1}) == {"a": {"x": 1}, "b": {}} + + +def test_a_value_that_contains_itself_is_a_config_error(): + loop: dict[str, Any] = {} + loop["b"] = loop + with pytest.raises(ConfigError, match="refers to itself"): + apply_overrides({"a": loop}, {}) + with pytest.raises(ConfigError, match="refers to itself"): + apply_overrides({}, {"a": loop}) + + +# --------------------------------------------------------------------------- +# Composed with the base: the override beats the file, which beats the +# defaults, and every schema error is pydantic's at the path the override named. + + +def test_an_override_beats_the_file_which_beats_the_defaults(tmp_path): + config = load(tmp_path, ["--model.num_interactions", "3"]) + assert config.model.num_interactions == 3 # the override beats the file's 4 + assert config.model.radial.cutoff == 4.5 # the file's other values survive + assert config.model.radial.num_bessel == 8 # defaults fill the rest + assert load(tmp_path, []) == DemoConfig.from_dict(FILE_VALUES) + + +def test_values_are_handed_to_pydantic_as_given(tmp_path): + argv = ["--seed=9", "--data.train_file", "null", "--data.heads", HEADS_JSON] + config = load(tmp_path, argv, {}) + assert config.seed == 9 # pydantic's lax coercion + assert config.data.train_file is None + assert config.data.heads == ["a", "b"] + assert load(tmp_path, ["--stage_two.start_epoch", "50"], {}).stage_two == ( + StageTwoSection(start_epoch=50) + ) + + +def test_a_value_of_the_wrong_type_is_a_validation_error(tmp_path): + with pytest.raises(ValidationError, match="seed"): + load(tmp_path, ["--seed", "seven"]) + + +@pytest.mark.parametrize( + "dotted_path", ["nmae", "model.num_interaction", "stage_two.start"] +) +def test_an_unknown_key_in_an_override_is_reported_at_its_path(tmp_path, dotted_path): + with pytest.raises(ValidationError) as excinfo: + load(tmp_path, [f"--{dotted_path}", "1"], {}) + assert error_locations(excinfo) == [dotted_path] + assert excinfo.value.errors()[0]["type"] == "extra_forbidden" + + +def test_an_override_writes_into_a_kinds_field(tmp_path): + config = load(tmp_path, ["--loss.delta", "2"], {"loss": {"kind": "huber"}}) + assert config.loss == Huber(kind="huber", delta=2.0) + config = load(tmp_path, ["--loss", json.dumps({"kind": "huber", "delta": 2})], {}) + assert config.loss == Huber(kind="huber", delta=2.0) + assert config.to_user_dict() == {"loss": {"kind": "huber", "delta": 2.0}} + + +def test_an_override_that_changes_the_kind_leaves_the_old_settings_to_pydantic( + tmp_path, +): + # The file's `delta` stays in the dict and is an unknown key of the new kind. + huber = {"loss": {"kind": "huber", "delta": 0.5}} + with pytest.raises(ValidationError) as excinfo: + load(tmp_path, ["--loss.kind", "weighted"], huber) + assert error_locations(excinfo) == ["loss.weighted.delta"] diff --git a/packages/mace-core/tests/test_mace_core_config_kinds.py b/packages/mace-core/tests/test_mace_core_config_kinds.py index 63301c20d..36c5c0c1a 100644 --- a/packages/mace-core/tests/test_mace_core_config_kinds.py +++ b/packages/mace-core/tests/test_mace_core_config_kinds.py @@ -1,26 +1,16 @@ -"""A field of several kinds of section (a discriminated union) under the file -and dotted-override contract. A config writes each kind's settings under the -kind's name (`loss: {huber: {delta: 0.1}}`, `--loss.huber.delta 0.1`); the -settings of every kind are kept across files and overrides. Which kind runs is -selected by `kind: huber`, `--loss.kind huber` or the bare name `--loss huber`; -a single kind key selects itself; the last selection wins. An override that -changed nothing about the config that runs warns, a file never warns; two kinds -with no selection is an error. `kind` is the tag only under a kinds field: a -plain section field of the same class keeps it as a key. Code sees the union.""" +"""A field of several kinds of section (a discriminated union on `kind`) is an +ordinary pydantic feature under the one-file contract: a config +writes it in the tagged form (`loss: {kind: huber, delta: 0.1}`), pydantic +picks the variant by the tag and reports every error under it +(`loss.huber.delta`), and both exports write the tag back so the dict loads to +the same config. Nothing in the loader knows about kinds.""" import json -import re -import warnings -from typing import Annotated, Any, Literal +from typing import Annotated, Literal import pytest -from mace_core.config import ( - ConfigError, - ConfigSection, - ConfigWarning, - ReforgeBaseConfig, -) -from pydantic import BaseModel, Field, ValidationError +from mace_core.config import ConfigSection, ReforgeBaseConfig +from pydantic import Field, ValidationError #: A warning the test did not ask for is a failure. pytestmark = pytest.mark.filterwarnings("error") @@ -57,10 +47,6 @@ class Universal(ConfigSection): huber_delta: float = 0.01 -class Plain(ConfigSection): - p: int = 0 - - Choice = Annotated[Weighted | Huber | Universal, Field(discriminator="kind")] @@ -83,7 +69,6 @@ class LossConfig(ReforgeBaseConfig): energy_weight: float = 1.0 choice: Choice = Weighted() opt: OptChoice = NoChoice() - heads: dict[str, Plain] = Field(default_factory=dict) per_head: dict[str, HeadSection] = Field(default_factory=dict) layers: list[HeadSection] = Field(default_factory=list) @@ -100,50 +85,6 @@ class RequiredConfig(ReforgeBaseConfig): ) -# --------------------------------------------------------------------------- -# Expectations. A row is (file values, command line, expectation); the -# expectation is a predicate on the loaded config, optionally with the -# warnings the load must emit, or an error spec. - -KINDS = "weighted, huber, universal" -OPT_KINDS = "weighted, huber, universal, none" - - -class Raises: - def __init__(self, error_type, *fragments): - self.error_type = error_type - self.fragments = fragments - - -class Warns: - """A predicate plus the exact warning messages, in order.""" - - def __init__(self, predicate, *messages): - self.predicate = predicate - self.messages = messages - - -def error(*fragments): - """A `ConfigError` whose message holds every fragment.""" - return Raises(ConfigError, *fragments) - - -def no_effect(source, field, kind, running): - return f"{source}: {field}.{kind} has no effect, {field} runs {running}" - - -def overridden(source, field, running): - return f"{source} is overridden: {field} runs {running}" - - -def needs_kind(field, *kinds): - """The kinds error up to the sources; the first kind is the example.""" - choices = " or ".join(f"kind: {kind}" for kind in kinds) - return ( - f"{field} needs a kind; write {choices} in a file, or pass --{field} {kinds[0]}" - ) - - def huber(delta=0.01, sub=SubX, **sub_fields): return lambda c: ( isinstance(c.choice, Huber) @@ -153,604 +94,124 @@ def huber(delta=0.01, sub=SubX, **sub_fields): ) -def weighted(stress_weight=0.0): - return lambda c: ( - isinstance(c.choice, Weighted) and c.choice.stress_weight == stress_weight - ) - - -def universal(huber_delta=0.01): - return lambda c: ( - isinstance(c.choice, Universal) and c.choice.huber_delta == huber_delta - ) - - -def opt(kind, **fields): - return lambda c: ( - isinstance(c.opt, kind) - and all(getattr(c.opt, k) == v for k, v in fields.items()) - ) +def load(tmp_path, root, file_values): + path = tmp_path / "config.json" + path.write_text(json.dumps(file_values)) + return root.load(path) -def per_head(name, kind, **fields): - return lambda c: ( - isinstance(c.per_head[name].loss, kind) - and all(getattr(c.per_head[name].loss, k) == v for k, v in fields.items()) - ) +def error_locations(excinfo): + """The dotted location of every error a `ValidationError` carries.""" + return [".".join(map(str, error["loc"])) for error in excinfo.value.errors()] -def layer(index, kind, **fields): - return lambda c: ( - isinstance(c.layers[index].loss, kind) - and all(getattr(c.layers[index].loss, k) == v for k, v in fields.items()) - ) +# --------------------------------------------------------------------------- +# Loading: the tagged form, and pydantic's errors under the tag. -HUBER_FILE = {"choice": {"huber": {"delta": 0.5}}} - -SWITCHING = { - "file kind, no cli": (HUBER_FILE, "", huber(0.5)), - "cli name selects; the file's settings of the other kind stay unused": ( - HUBER_FILE, - "--choice weighted", - weighted(), - ), - "cli key of a second kind without a selection is an error": ( - HUBER_FILE, - "--choice.weighted.stress_weight 5", - error( - needs_kind("choice", "huber", "weighted"), - "; huber from ", - "config.json, weighted from --choice.weighted.stress_weight 5", - ), - ), - "cli json of a second kind without a selection is an error": ( - HUBER_FILE, - '--choice {"universal": {"huber_delta": 3}}', - error( - needs_kind("choice", "huber", "universal"), - 'config.json, universal from --choice {"universal": {"huber_delta": 3}}', - ), - ), - "the tag alone selects on the cli": ({}, '--choice {"kind": "huber"}', huber()), - "the tag beside the settings selects in the file": ( - {"choice": {"kind": "weighted", "huber": {"delta": 0.5}, "weighted": {}}}, - "", - weighted(), - ), - "cli name of the file's kind keeps the file's keys": ( - HUBER_FILE, - "--choice huber", - huber(0.5), - ), - "cli key of the file's kind merges": ( - HUBER_FILE, - "--choice.huber.sub y", - huber(0.5, SubY), - ), - "a selection, then a key of another kind: the key warns": ( - HUBER_FILE, - "--choice weighted --choice.huber.sub y", - Warns( - weighted(), - no_effect("--choice.huber.sub y", "choice", "huber", "weighted"), - ), - ), - "keys of two kinds on the cli without a selection is an error": ( - {}, - "--choice.huber.delta 2 --choice.weighted.stress_weight 5", - error( - needs_kind("choice", "huber", "weighted") - + "; huber from --choice.huber.delta 2, " - "weighted from --choice.weighted.stress_weight 5" - ), - ), - "keys of two kinds on the cli, then a selection": ( - {}, - "--choice.huber.delta 2 --choice.weighted.stress_weight 5 --choice huber", - Warns( - huber(2.0), - no_effect( - "--choice.weighted.stress_weight 5", "choice", "weighted", "huber" - ), - ), - ), - "no file, default kind": ({}, "", weighted()), - "no file, cli key of another kind": ({}, "--choice.huber.delta 2", huber(2.0)), - "empty section is the default kind": ({"choice": {}}, "", weighted()), - "empty section with cli key": ( - {"choice": {}}, - "--choice.huber.delta 2", - huber(2.0), - ), - "bare name in the file": ({"choice": "huber"}, "", huber()), - "bare name in the file, cli key of that kind": ( - {"choice": "huber"}, - "--choice.huber.delta 2", - huber(2.0), - ), - "bare name in the file, cli key of another kind warns": ( - {"choice": "huber"}, - "--choice.weighted.stress_weight 5", - Warns( - huber(), - no_effect( - "--choice.weighted.stress_weight 5", "choice", "weighted", "huber" - ), - ), - ), - "= form": (HUBER_FILE, "--choice.huber.delta=2", huber(2.0)), - "three selections in a row: the last runs, the lost cli one warns": ( - HUBER_FILE, - "--choice weighted --choice universal", - Warns(universal(), overridden("--choice weighted", "choice", "universal")), - ), -} - -ERRORS = { - "two kinds in the file": ( - {"choice": {"huber": {}, "weighted": {}}}, - "", - error( - needs_kind("choice", "huber", "weighted"), - "; huber from ", - "config.json, weighted from ", - ), - ), - "two kinds in one json override": ( - {}, - '--choice {"huber": {}, "weighted": {}}', - error( - needs_kind("choice", "huber", "weighted") - + '; huber from --choice {"huber": {}, "weighted": {}}, ' - 'weighted from --choice {"huber": {}, "weighted": {}}' - ), - ), - "unknown kind in the file": ( - {"choice": {"hubr": {}}}, - "", - error( - "unknown config key 'choice.hubr'; did you mean 'choice.huber'?; " - f"the kinds of choice are {KINDS}" - ), - ), - "unknown bare name in the file": ( - {"choice": "hubr"}, - "", - error("unknown config key 'choice.hubr'; did you mean 'choice.huber'?"), - ), - "unknown kind on the cli": ( - {}, - "--choice hubr", - error("unknown config key 'choice.hubr'; did you mean 'choice.huber'?"), - ), - "unknown dotted kind on the cli": ( - {}, - "--choice.hubr.delta 2", - error( - "unknown config key 'choice.hubr'; did you mean 'choice.huber'?; " - f"the kinds of choice are {KINDS}" - ), - ), - "a setting beside the kinds is unknown and the kinds are listed": ( - {}, - "--choice.delta 2", - error(f"unknown config key 'choice.delta'; the kinds of choice are {KINDS}"), - ), - "key of another kind in the file": ( - {"choice": {"weighted": {"delta": 2}}}, - "", - error("unknown config key 'choice.weighted.delta'"), - ), - "the tag is not a key, on the cli": ( - {}, - "--choice.huber.kind huber", - error( - "unknown config key 'choice.huber.kind'; the key huber already names " - "the kind" - ), - ), - "the tag is not a key, in the file": ( - {"choice": {"huber": {"kind": "weighted"}}}, - "", - error( - "unknown config key 'choice.huber.kind'; the key huber already names " - "the kind" - ), - ), - "the flat form fails on the setting beside the tag": ( - {"choice": {"kind": "huber", "delta": 2}}, - "", - error(f"unknown config key 'choice.delta'; the kinds of choice are {KINDS}"), - ), - "a selection that is not a kind": ( - {}, - "--choice.kind hubr", - error( - "unknown config key 'choice.hubr'; did you mean 'choice.huber'?; " - f"the kinds of choice are {KINDS}" - ), - ), - "a list at a kinds field": ( - {"choice": [1]}, - "", - error( - "choice must be the name of a kind or a mapping under one, one of " - f"{KINDS}; got [1]" - ), - ), - "a number at a kinds field": ( - {"choice": 3}, - "", - error( - "choice must be the name of a kind or a mapping under one, one of " - f"{KINDS}; got 3" - ), - ), - "a scalar under a kind": ( - {"choice": {"huber": 3}}, - "", - error("choice.huber must be a mapping of its keys; got 3"), - ), - "null at a kinds field": ( - {"choice": None}, - "", - error( - "choice must be the name of a kind or a mapping under one, one of " - f"{KINDS}; got null" - ), - ), - "settings of a kind that does not run are still checked": ( - {"choice": {"huber": {"nonsense": 1}}}, - "--choice weighted", - error("unknown config key 'choice.huber.nonsense'"), - ), -} - -NULL = { - "null at a kinds field is an error even when a key follows": ( - HUBER_FILE, - "--choice null --choice.huber.sub y", - error( - "choice must be the name of a kind or a mapping under one, one of " - f"{KINDS}; got null" - ), - ), - "null under a kind, on the cli": ( - HUBER_FILE, - "--choice.huber null", - error("choice.huber must be a mapping of its keys; got null"), - ), - "null under an unknown kind is an unknown kind": ( - {"choice": {"hubr": None}}, - "", - error("unknown config key 'choice.hubr'; did you mean 'choice.huber'?"), - ), - "null under a kind, in the file": ( - {"choice": {"huber": {"delta": 0.5}, "weighted": None}}, - "", - error("choice.weighted must be a mapping of its keys; got null"), - ), -} - -NONE_KIND = { - "none kind, absent": ({}, "", opt(NoChoice)), - "none kind, by name": ({}, "--opt none", opt(NoChoice)), - "none kind, cli key": ({}, "--opt.huber.delta 2", opt(Huber, delta=2.0)), - "none kind, file null": ( - {"opt": None}, - "", - error( - "opt must be the name of a kind or a mapping under one, one of " - f"{OPT_KINDS}; got null" - ), - ), - "none kind, cli null after the file": ( - {"opt": {"huber": {}}}, - "--opt null", - error( - "opt must be the name of a kind or a mapping under one, one of " - f"{OPT_KINDS}; got null" - ), - ), - "none kind, back to none": ({"opt": {"huber": {}}}, "--opt none", opt(NoChoice)), -} - -NESTED = { - "nested kind by name": (HUBER_FILE, "--choice.huber.sub y", huber(0.5, SubY)), - "nested kind by key": ( - HUBER_FILE, - "--choice.huber.sub.y.b 7", - huber(0.5, SubY, b=7.0), - ), - "nested kind in the file, cli key of it": ( - {"choice": {"huber": {"sub": {"y": {"b": 3}}}}}, - "--choice.huber.sub.y.b 7", - huber(0.01, SubY, b=7.0), - ), - "nested selection runs; the file's settings of the other kind stay unused": ( - {"choice": {"huber": {"sub": {"y": {"b": 3}}}}}, - "--choice.huber.sub x", - huber(0.01, SubX), - ), - "nested key of another kind on the cli warns": ( - {"choice": {"huber": {"sub": {"kind": "y", "y": {"b": 3}}}}}, - "--choice.huber.sub.x.a 5", - Warns( - huber(0.01, SubY, b=3.0), - no_effect("--choice.huber.sub.x.a 5", "choice.huber.sub", "x", "y"), - ), - ), - "nested empty section is its default kind": ( - {"choice": {"huber": {"sub": {}}}}, - "", - huber(0.01, SubX), - ), - "outer selection: the nested settings stay unused": ( - {"choice": {"huber": {"sub": {"y": {"b": 3}}}}}, - "--choice weighted", - weighted(), - ), - "unknown nested kind": ( - {"choice": {"huber": {"sub": {"z": {}}}}}, - "", - error( - "unknown config key 'choice.huber.sub.z'", - "the kinds of choice.huber.sub are x, y", - ), - ), - "unknown key under a nested kind": ( - {"choice": {"huber": {"sub": {"y": {"a": 1}}}}}, - "", - error("unknown config key 'choice.huber.sub.y.a'"), - ), -} - -COLLECTIONS = { - "kind inside a dict value, by name": ( - {"per_head": {"h": {"loss": "huber"}}}, - "", - per_head("h", Huber), - ), - "kind inside a dict value, by key": ( - {"per_head": {"h": {"loss": {"huber": {"delta": 2}}}}}, - "", - per_head("h", Huber, delta=2.0), - ), - "kind inside a dict value, json merges the entry": ( - {"per_head": {"h": {"loss": {"huber": {"delta": 2}}}}}, - '--per_head {"h": {"loss": "huber"}}', - per_head("h", Huber, delta=2.0), - ), - "kind inside a dict value, json selects another kind": ( - {"per_head": {"h": {"loss": {"huber": {"delta": 2}}}}}, - '--per_head {"h": {"loss": "weighted"}}', - per_head("h", Weighted), - ), - "kind inside a dict value, dotted key of another kind warns": ( - {"per_head": {"h": {"loss": {"kind": "huber", "huber": {"delta": 2}}}}}, - "--per_head.h.loss.weighted.stress_weight 5", - Warns( - per_head("h", Huber, delta=2.0), - no_effect( - "--per_head.h.loss.weighted.stress_weight 5", - "per_head.h.loss", - "weighted", - "huber", - ), - ), - ), - "kind inside a dict value, empty is the default": ( - {"per_head": {"h": {"loss": {}}}}, - "", - per_head("h", Weighted), - ), - "unknown key inside a dict value names the kind": ( - {"per_head": {"h": {"loss": {"huber": {"stress_weight": 1}}}}}, - "", - error("unknown config key 'per_head.h.loss.huber.stress_weight'"), - ), - "kind inside a list item": ( - {"layers": [{"loss": "huber"}, {"loss": {"universal": {"huber_delta": 3}}}]}, - "", - lambda c: layer(0, Huber)(c) and layer(1, Universal, huber_delta=3.0)(c), - ), - "list replaced whole by json": ( - {"layers": [{"loss": "huber"}]}, - '--layers [{"loss": {"weighted": {"stress_weight": 5}}}]', - lambda c: len(c.layers) == 1 and layer(0, Weighted, stress_weight=5.0)(c), - ), - "unknown key inside a list item": ( - {"layers": [{"loss": {"huber": {"stress_weight": 1}}}]}, - "", - error("unknown config key 'layers.0.loss.huber.stress_weight'"), - ), - "unknown kind inside a list item": ( - {"layers": [{"loss": {"hubr": {}}}]}, - "", - error( - "unknown config key 'layers.0.loss.hubr'; did you mean " - f"'layers.0.loss.huber'?; the kinds of layers.0.loss are {KINDS}" - ), - ), - "list item cannot be addressed by a dotted key": ( - {}, - "--layers.0.loss huber", - error("unknown config key 'layers.0'; layers is written whole"), - ), -} - -WARNINGS = { - "a selection lost to a later one warns, the settings stay": ( - {}, - "--choice huber --choice.huber.delta 2 --choice weighted", - Warns( - weighted(), - overridden("--choice huber", "choice", "weighted"), - no_effect("--choice.huber.delta 2", "choice", "huber", "weighted"), - ), - ), - "a selection by the tag lost to a later one warns": ( - {}, - "--choice.kind huber --choice weighted", - Warns(weighted(), overridden("--choice.kind huber", "choice", "weighted")), - ), - "a json override that selects one kind and tunes another warns once": ( - {}, - '--choice {"kind": "weighted", "huber": {"delta": 1, "sub": "y"}}', - Warns( - weighted(), - '--choice {"kind": "weighted", "huber": {"delta": 1, "sub": "y"}}: ' - "choice.huber has no effect, choice runs weighted", +@pytest.mark.parametrize("extension", [".toml", ".yaml", ".json"]) +def test_kinds_load_from_every_file_format(tmp_path, extension): + text = { + ".toml": '[choice]\nkind = "huber"\ndelta = 0.5\n[choice.sub]\nkind = "y"\n' + 'b = 3\n[opt]\nkind = "universal"\n', + ".yaml": "choice:\n kind: huber\n delta: 0.5\n sub:\n kind: y\n" + " b: 3\nopt:\n kind: universal\n", + ".json": json.dumps( + { + "choice": {"kind": "huber", "delta": 0.5, "sub": {"kind": "y", "b": 3}}, + "opt": {"kind": "universal"}, + } ), - ), - "the same selection twice: the first is overridden": ( - {}, - "--choice huber --choice huber", - Warns(huber(), overridden("--choice huber", "choice", "huber")), - ), - "a setting of the running kind is silent": ( - {"choice": "huber"}, - "--choice.huber.delta 3", - huber(3.0), - ), - "an empty mapping before a setting is silent": ( - {}, - "--choice {} --choice.huber.delta 1", - huber(1.0), - ), - "a file tuning several kinds never warns": ( - {"choice": {"kind": "huber", "huber": {"delta": 0.5}, "weighted": {}}}, - "--energy_weight 2", - lambda c: c.energy_weight == 2.0 and huber(0.5)(c), - ), -} - -OTHER = { - "keys of other fields are untouched by a selection": ( - {"energy_weight": 3.0, **HUBER_FILE}, - "--choice weighted", - lambda c: c.energy_weight == 3.0 and weighted()(c), - ), - "a required key of the kind supplied by the cli": ( - {"choice": {"huber": {}}}, - "--choice.huber.path p", - lambda c: isinstance(c.choice, HuberRequired) and c.choice.path == "p", - ), - "a required key of the kind missing": ( - {"choice": "huber"}, - "", - Raises(ValidationError, "choice.huber.path", "Field required"), - ), -} - -ROWS = { - **SWITCHING, - **ERRORS, - **NULL, - **NONE_KIND, - **NESTED, - **COLLECTIONS, - **WARNINGS, -} - - -def split_cli(cli): - """`--a.b 1 --c {"d": 2}` as argv: a JSON value keeps its spaces.""" - argv = [] - for chunk in filter(None, re.split(r"\s+(?=--)", cli)): - argv.extend(chunk.split(" ", 1)) - return argv - - -def load(tmp_path, root, file_values, cli): - path = tmp_path / "config.json" - path.write_text(json.dumps(file_values)) - return root.load(path, cli_overrides=split_cli(cli)) - - -def check(tmp_path, root, file_values, cli, expected): - if isinstance(expected, Raises): - with pytest.raises(expected.error_type) as info, warnings.catch_warnings(): - warnings.simplefilter("ignore", ConfigWarning) - load(tmp_path, root, file_values, cli) - for fragment in expected.fragments: - assert fragment in str(info.value), str(info.value) - return - predicate, messages = expected, () - if isinstance(expected, Warns): - predicate, messages = expected.predicate, expected.messages - with warnings.catch_warnings(record=True) as caught: - warnings.simplefilter("always") - config = load(tmp_path, root, file_values, cli) - assert [str(w.message) for w in caught] == list(messages) - assert all(issubclass(w.category, ConfigWarning) for w in caught) - assert predicate(config), config + }[extension] + path = tmp_path / f"config{extension}" + path.write_text(text) + config = LossConfig.load(path) + assert huber(0.5, SubY, b=3.0)(config) and isinstance(config.opt, Universal) -@pytest.mark.parametrize("row", ROWS, ids=ROWS) -def test_kind_selection(tmp_path, row): - check(tmp_path, LossConfig, *ROWS[row]) +@pytest.mark.parametrize("value", [{"delta": 0.5}, {}], ids=["settings", "empty"]) +def test_a_kinds_field_without_its_tag_is_a_validation_error(tmp_path, value): + # Even though the field has a default: pydantic needs the tag to pick. + with pytest.raises(ValidationError) as excinfo: + load(tmp_path, LossConfig, {"choice": value}) + assert error_locations(excinfo) == ["choice"] + assert excinfo.value.errors()[0]["type"] == "union_tag_not_found" -@pytest.mark.parametrize("row", OTHER, ids=OTHER) -def test_kind_selection_on_other_roots(tmp_path, row): - file_values, cli, expected = OTHER[row] - root = RequiredConfig if "required" in row else LossConfig - check(tmp_path, root, file_values, cli, expected) +def test_a_wrong_tag_is_a_validation_error_listing_the_kinds(tmp_path): + with pytest.raises(ValidationError) as excinfo: + load(tmp_path, LossConfig, {"choice": {"kind": "hubr"}}) + assert error_locations(excinfo) == ["choice"] + (error,) = excinfo.value.errors() + assert error["type"] == "union_tag_invalid" + assert all(kind in error["msg"] for kind in ("weighted", "huber", "universal")) @pytest.mark.parametrize( - ("extension", "text", "shown"), + ("file_values", "location"), [ - (".toml", "choice = 2020-01-01\n", '"2020-01-01"'), - (".yaml", "choice:\n huber: 2020-01-01\n", '"2020-01-01"'), - (".yaml", "choice: !!set {huber: null}\n", "\"{'huber'}\""), + ( + {"choice": {"kind": "huber", "stress_weight": 1}}, + "choice.huber.stress_weight", + ), + ( + {"choice": {"kind": "huber", "sub": {"kind": "y", "a": 1}}}, + "choice.huber.sub.y.a", + ), + ( + {"per_head": {"h": {"loss": {"kind": "huber", "stress_weight": 1}}}}, + "per_head.h.loss.huber.stress_weight", + ), + ( + {"layers": [{"loss": {"kind": "huber", "stress_weight": 1}}]}, + "layers.0.loss.huber.stress_weight", + ), ], + ids=["top", "nested", "dict value", "list item"], ) -def test_a_value_json_cannot_show_is_still_a_config_error( - tmp_path, extension, text, shown +def test_unknown_keys_under_a_kind_are_reported_under_the_tag( + tmp_path, file_values, location ): - path = tmp_path / f"config{extension}" - path.write_text(text) - with pytest.raises(ConfigError, match=rf"got {shown}$"): - LossConfig.load(path) + with pytest.raises(ValidationError) as excinfo: + load(tmp_path, LossConfig, file_values) + assert error_locations(excinfo) == [location] + assert excinfo.value.errors()[0]["type"] == "extra_forbidden" -@pytest.mark.parametrize("extension", [".toml", ".yaml"]) -def test_kinds_load_from_every_file_format(tmp_path, extension): - text = { - ".toml": "[choice.huber]\ndelta = 0.5\n[choice.huber.sub.y]\nb = 3\n" - "[opt]\nuniversal = {}\n", - ".yaml": "choice:\n huber:\n delta: 0.5\n sub:\n y:\n" - " b: 3\nopt: universal\n", - }[extension] - path = tmp_path / f"config{extension}" - path.write_text(text) - config = LossConfig.load(path) - assert huber(0.5, SubY, b=3.0)(config) and isinstance(config.opt, Universal) +def test_a_required_key_of_the_kind_is_reported_under_the_tag(tmp_path): + with pytest.raises(ValidationError) as excinfo: + load(tmp_path, RequiredConfig, {"choice": {"kind": "huber"}}) + assert error_locations(excinfo) == ["choice.huber.path"] + assert excinfo.value.errors()[0]["type"] == "missing" + config = load(tmp_path, RequiredConfig, {"choice": {"kind": "huber", "path": "p"}}) + assert isinstance(config.choice, HuberRequired) and config.choice.path == "p" # --------------------------------------------------------------------------- -# The dicts: kind as key, a fixed point, and only what was set. +# The dicts: the tag is written back, a fixed point, and only what was set. @pytest.mark.parametrize( - "cli", + "file_values", [ - "", - "--choice huber", - "--choice.huber.sub.y.b 7 --opt universal", - '--per_head {"h": {"loss": "huber"}} --layers [{"loss": {"universal": {}}}]', + {}, + {"choice": {"kind": "huber"}}, + { + "choice": {"kind": "huber", "sub": {"kind": "y", "b": 7}}, + "opt": {"kind": "universal"}, + }, + { + "per_head": {"h": {"loss": {"kind": "huber"}}}, + "layers": [{"loss": {"kind": "universal"}}], + }, ], + ids=["defaults", "huber", "nested", "collections"], ) -def test_resolved_dict_is_a_fixed_point_and_writes_the_kind_as_key(tmp_path, cli): - config = load(tmp_path, LossConfig, {}, cli) +def test_resolved_dict_is_a_fixed_point_and_writes_the_tag(tmp_path, file_values): + config = load(tmp_path, LossConfig, file_values) resolved = config.to_resolved_dict() - kind = type(config.choice).model_fields["kind"].default - assert list(resolved["choice"]) == [kind] - assert "kind" not in resolved["choice"][kind] - reloaded = load(tmp_path, LossConfig, resolved, "") + assert resolved["choice"]["kind"] == config.choice.kind + reloaded = load(tmp_path, LossConfig, resolved) assert reloaded.to_resolved_dict() == resolved assert reloaded == config @@ -758,43 +219,28 @@ def test_resolved_dict_is_a_fixed_point_and_writes_the_kind_as_key(tmp_path, cli def test_resolved_dict_of_the_defaults(): assert LossConfig().to_resolved_dict() == { "energy_weight": 1.0, - "choice": {"weighted": {"stress_weight": 0.0}}, - "opt": {"none": {}}, - "heads": {}, + "choice": {"kind": "weighted", "stress_weight": 0.0}, + "opt": {"kind": "none"}, "per_head": {}, "layers": [], } -def test_user_dict_holds_the_kind_and_only_what_was_set(tmp_path): - config = load(tmp_path, LossConfig, {"choice": "huber"}, "--choice.huber.delta 2") - assert config.to_user_dict() == {"choice": {"huber": {"delta": 2.0}}} - # A kind chosen by default is what ran, so it is set. - assert load(tmp_path, LossConfig, {"choice": {}}, "").to_user_dict() == { - "choice": {"weighted": {}} - } - # Built in code, the tag is an unset default; the kind is still the key. +def test_user_dict_holds_the_tag_and_only_what_was_set(tmp_path): + config = load(tmp_path, LossConfig, {"choice": {"kind": "huber", "delta": 2.0}}) + assert config.to_user_dict() == {"choice": {"kind": "huber", "delta": 2.0}} + assert load(tmp_path, LossConfig, config.to_user_dict()) == config + # The documented caveat: built in code without its tag, a variant exports + # without `kind` and the dict does not load back; pass the tag, or use + # `to_resolved_dict`. assert LossConfig(choice=Huber(delta=2)).to_user_dict() == { - "choice": {"huber": {"delta": 2.0}} - } - - -def test_the_empty_mapping_keeps_the_default_instance_with_its_settings(tmp_path): - class Tuned(ReforgeBaseConfig): - choice: Choice = Weighted(stress_weight=2.0) - made: Choice = Field(default_factory=lambda: Huber(delta=9.0, sub=SubY(b=4.0))) - - config = load(tmp_path, Tuned, {"choice": {}, "made": {}}, "") - assert config == Tuned() - assert config.to_user_dict() == { - "choice": {"weighted": {"stress_weight": 2.0}}, - "made": {"huber": {"delta": 9.0, "sub": {"y": {"b": 4.0}}}}, + "choice": {"delta": 2.0} } - assert load(tmp_path, Tuned, config.to_user_dict(), "") == config - # A kind key with no settings runs the class defaults, not the instance's. - assert load(tmp_path, Tuned, {"choice": {"weighted": {}}}, "").choice == Weighted() - tuned = load(tmp_path, Tuned, {"choice": {}}, "--choice.weighted.stress_weight 3") - assert tuned.choice == Weighted(stress_weight=3.0) + tagged = LossConfig(choice=Huber(kind="huber", delta=2)) + assert tagged.to_user_dict() == {"choice": {"kind": "huber", "delta": 2.0}} + assert LossConfig(choice=Huber(delta=2)).to_resolved_dict()["choice"]["kind"] == ( + "huber" + ) def test_json_schema_is_produced_in_both_modes(): @@ -803,215 +249,16 @@ def test_json_schema_is_produced_in_both_modes(): assert set(schema["properties"]) == set(LossConfig.model_fields), mode -def test_code_sees_the_union_and_may_construct_it_either_way(): - by_key = {"choice": {"huber": {"delta": 2}}} - by_tag = {"choice": {"kind": "huber", "delta": 2}} - assert LossConfig(choice=Huber(delta=2)).choice == Huber(delta=2) - assert LossConfig.model_validate(by_key).choice == Huber(delta=2) - assert LossConfig.model_validate(by_tag).choice == Huber(delta=2) - - def test_a_variant_class_as_a_plain_field_keeps_kind_as_a_key(tmp_path): class Reuse(ReforgeBaseConfig): direct: Huber = Huber() - resolved = {"direct": {"kind": "huber", "delta": 0.01, "sub": {"x": {"a": 1.0}}}} + resolved = { + "direct": {"kind": "huber", "delta": 0.01, "sub": {"kind": "x", "a": 1.0}} + } assert Reuse().to_resolved_dict() == resolved - assert load(tmp_path, Reuse, resolved, "").to_resolved_dict() == resolved - config = load(tmp_path, Reuse, {}, "--direct.kind huber --direct.sub y") + assert load(tmp_path, Reuse, resolved).to_resolved_dict() == resolved + config = load(tmp_path, Reuse, {"direct": {"kind": "huber", "sub": {"kind": "y"}}}) assert config.direct == Huber(sub=SubY()) with pytest.raises(ValidationError, match=r"direct\.kind"): - load(tmp_path, Reuse, {}, "--direct.kind weighted") - - -# --------------------------------------------------------------------------- -# Schema rules. - - -class NotASection(BaseModel): - kind: Literal["plain"] = "plain" - - -def test_union_shapes_the_contract_cannot_keep_are_rejected_at_class_definition(): - class IntTag(ConfigSection): - kind: Literal[1] = 1 - - class TwoTags(ConfigSection): - kind: Literal["a", "b"] = "a" - - class NamedLikeTheTag(ConfigSection): - kind: Literal["kind"] = "kind" - - class Colliding(ConfigSection): - """In the flat form `{kind: c, weighted: 9}`, the field would read as - the settings of the sibling kind.""" - - kind: Literal["c"] = "c" - weighted: float = 3.0 - - shapes = { - r"choice defaults to None, which is not a kind; default to a variant": ( - Choice, - None, - ), - r"choice admits None, which is not a kind; default to a variant, or for none": ( - Choice | None, - Weighted(), - ), - r"choice has variant NamedLikeTheTag whose kind is named 'kind' like the tag": ( - Annotated[Weighted | NamedLikeTheTag, Field(discriminator="kind")], - Weighted(), - ), - r"choice is a union of sections without a discriminator": ( - Weighted | Huber, - Weighted(), - ), - r"choice is a union of sections inside a dict, list or tuple": ( - list[Choice], - [], - ), - r"choice holds NotASection, which is not a ConfigSection": ( - Annotated[Weighted | NotASection, Field(discriminator="kind")], - Weighted(), - ), - r"choice has variant IntTag whose kind must be a Literal of exactly one": ( - Annotated[Weighted | IntTag, Field(discriminator="kind")], - Weighted(), - ), - r"choice has variant TwoTags whose kind must be a Literal of exactly one": ( - Annotated[Weighted | TwoTags, Field(discriminator="kind")], - Weighted(), - ), - r"choice has a default that is not one of its variants; write e.g. Weighted": ( - Choice, - Plain(), - ), - r"choice has variant Colliding with a field named like the kind weighted; ": ( - Annotated[Weighted | Colliding, Field(discriminator="kind")], - Weighted(), - ), - r"choice mixes its kinds with dict; a kinds field holds its variants only": ( - Choice | dict[str, int], - Weighted(), - ), - } - for message, (annotation, default) in shapes.items(): - with pytest.raises(TypeError, match=message): - type( - "Bad", - (ConfigSection,), - {"__annotations__": {"choice": annotation}, "choice": default}, - ) - - -def test_a_required_kinds_field_is_accepted(): - factory_calls = [] - - class Required(ReforgeBaseConfig): - choice: Choice - anything: Any = None - made: int = Field(default_factory=lambda: factory_calls.append(1) or 1) - - assert factory_calls == [], "the schema check must not run default factories" - with pytest.raises(ConfigError) as excinfo: - Required.load(cli_overrides=["--choice", "{}"]) - # Nothing wrote a kind, so no source is named. - assert str(excinfo.value) == ( - "choice needs a kind; write kind: weighted or kind: huber or kind: universal " - "in a file, or pass --choice weighted" - ) - assert isinstance(Required.load(cli_overrides=["--choice", "huber"]).choice, Huber) - - -def test_a_default_factory_that_returns_no_variant_is_reported(tmp_path): - class Broken(ReforgeBaseConfig): - # The wrong result is the point; ty sees only the factory type. - choice: Choice = Field(default_factory=lambda: None) # ty: ignore[invalid-assignment] - - with pytest.raises( - TypeError, match=r"Broken\.choice has a default factory whose result is not" - ): - load(tmp_path, Broken, {"choice": {}}, "") - - -def test_several_kinds_inside_a_list_item_names_only_the_file_fix(tmp_path): - file_values = {"layers": [{"loss": {"huber": {}, "weighted": {}}}]} - with pytest.raises(ConfigError) as excinfo: - load(tmp_path, LossConfig, file_values, "") - path = tmp_path / "config.json" - assert str(excinfo.value) == ( - "layers.0.loss needs a kind; write kind: huber or kind: weighted in a file; " - f"huber from {path}, weighted from {path}" - ) - - -def test_several_files_each_writing_one_kind_are_named(tmp_path): - defaults = tmp_path / "defaults.json" - defaults.write_text(json.dumps({"choice": {"huber": {"delta": 2}}})) - user = tmp_path / "user.json" - user.write_text(json.dumps({"choice": {"weighted": {}}})) - with pytest.raises(ConfigError) as excinfo: - LossConfig.load([defaults, user]) - assert str(excinfo.value) == ( - "choice needs a kind; write kind: huber or kind: weighted in a file, or " - f"pass --choice huber; huber from {defaults}, weighted from {user}" - ) - - -def test_the_tag_inside_a_kind_is_an_error_under_model_validate(): - # Not a silent switch to the other kind: the tag is not a key under a kind. - message = r"choice\.huber\.kind is not a key; huber already names the kind" - with pytest.raises(ValidationError, match=message): - LossConfig.model_validate({"choice": {"huber": {"kind": "weighted"}}}) - with pytest.raises(ValidationError, match=message): - LossConfig.model_validate( - {"choice": {"kind": "huber", "huber": {"kind": "huber", "delta": 2}}} - ) - - -def test_a_subclass_of_a_variant_exports_under_the_variant_kind(tmp_path): - class SubHuber(Huber): - extra_knob: int = 5 - - config = LossConfig(choice=SubHuber(delta=3.0)) - resolved = config.to_resolved_dict() - assert resolved["choice"] == {"huber": {"delta": 3.0, "sub": {"x": {"a": 1.0}}}} - assert config.to_user_dict() == {"choice": {"huber": {"delta": 3.0}}} - reloaded = load(tmp_path, LossConfig, resolved, "") - assert reloaded.choice == Huber(delta=3.0) - assert reloaded.to_resolved_dict() == resolved - - -def test_a_kind_named_like_a_class_still_reads_as_a_kind_in_error_paths(): - class Root(ConfigSection): - kind: Literal["Root"] = "Root" - q: int = 0 - - class Config(ReforgeBaseConfig): - choice: Annotated[Weighted | Root, Field(discriminator="kind")] = Weighted() - - with pytest.raises( - ConfigError, match=r"'choice\.Root\.qq'; did you mean 'choice\.Root\.q'" - ): - Config.load(cli_overrides=["--choice", '{"Root": {"qq": 1}}']) - - -def test_a_collection_key_named_like_the_element_class_stays_in_error_paths(): - with pytest.raises( - ConfigError, match=r"'heads\.Plain\.pp'; did you mean 'heads\.Plain\.p'" - ): - LossConfig.load(cli_overrides=["--heads", '{"Plain": {"pp": 1}}']) - - -def test_warning_can_be_turned_into_an_error(tmp_path): - with warnings.catch_warnings(): - warnings.simplefilter("error", ConfigWarning) - with pytest.raises( - ConfigWarning, match=r"choice\.weighted has no effect, choice runs huber" - ): - load( - tmp_path, - LossConfig, - {"choice": "huber"}, - "--choice.weighted.stress_weight 5", - ) + load(tmp_path, Reuse, {"direct": {"kind": "weighted"}}) From a666d442d934d677abd1aca7c54387a3f59da05a Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Thu, 24 Sep 2026 12:45:32 +0200 Subject: [PATCH 09/14] Rewrite the config base as one file and one validation, the command line beside it (CORE-2 follow-up, #1556) `base.py` knows nothing of a command line. `load(config_file)` reads one TOML, YAML or JSON file (`read_config_file`) and validates it once (`from_dict`, the same validation for a caller that edits the parsed dict first; it does not write into the dict). Everything the schema rejects is pydantic's `ValidationError` at its dotted location, unknown keys included; `ConfigError` is raised only for a file that cannot be read or parsed. The two exports are the plain `model_dump` calls, with inf and nan kept as JSON constants. A kinds field is pydantic's discriminated union in its tagged form; the module has no knowledge of it. `cli.py` is the command-line half, importing only `ConfigError` from the base. `parse_overrides` reads `--a.b value` and `--a.b=value` tokens into a mapping of dotted path to value (`null`, `[` and `{` values are JSON, anything else a string for pydantic to coerce; a token that is not an option, a missing value or bad JSON is a `ConfigError`). `apply_overrides` returns a copy of the parsed file with each override written at its path, creating mappings on the way, replacing a parent that is not a mapping, merging a mapping value into a mapping already there and replacing otherwise. The copy takes dicts and lists apart, so a YAML anchor's two keys stop sharing one object and an anchor that contains itself is a `ConfigError`. A command line builds its config as `Config.from_dict(apply_overrides(read_config_file(path), parse_overrides(rest)))`. `ConfigSection` keeps the definition-time checks only: `extra="forbid"` cannot be reopened, every model a field reaches is a section, no set of any kind (abstract ones included, since pydantic validates them to a frozenset), no aliases, excluded or computed fields. Every `from_dict` walks the sections its root reaches, rebuilds one a forward reference left incomplete, and checks each; re-checking is a few attribute reads per class and needs no record of what was checked. No validator, serializer or schema walk runs at load time. `_schema_rules.py` is deleted and `ConfigWarning` no longer exists; the public names are `ReforgeBaseConfig`, `ConfigSection`, `ConfigError` and `read_config_file` from the base, `parse_overrides` and `apply_overrides` from `cli`. The package README describes the new surface. --- packages/mace-core/README.md | 14 +- .../src/mace_core/config/__init__.py | 10 +- .../src/mace_core/config/_schema_rules.py | 240 ------ .../mace-core/src/mace_core/config/base.py | 706 ++++-------------- .../mace-core/src/mace_core/config/cli.py | 108 +++ 5 files changed, 265 insertions(+), 813 deletions(-) delete mode 100644 packages/mace-core/src/mace_core/config/_schema_rules.py create mode 100644 packages/mace-core/src/mace_core/config/cli.py diff --git a/packages/mace-core/README.md b/packages/mace-core/README.md index 83d2b013d..6039c2c0e 100644 --- a/packages/mace-core/README.md +++ b/packages/mace-core/README.md @@ -10,11 +10,15 @@ Distribution `mace-core`, import name `mace_core`. What is here so far: - `mace_core.config` — `ReforgeBaseConfig`, the pydantic base every v1 - config schema derives from: one TOML/YAML/JSON file plus dotted CLI - overrides (`--model.num_interactions 3`), - precedence defaults < file < CLI, - unknown keys rejected with the nearest valid neighbour named, and - `to_resolved_dict()` for the fully defaulted, round-trippable export. + config schema derives from: `load()` reads one TOML/YAML/JSON file and + validates it once (`from_dict()` for a caller that edits the parsed dict + first); unknown keys are pydantic errors at every level; + `to_resolved_dict()` is the fully defaulted, round-trippable export and + `to_user_dict()` holds only what was set. The base has no command-line + knowledge; `config.cli` holds the `--a.b value` grammar (`parse_overrides()` + to a mapping of dotted paths, `apply_overrides()` to write it into the + parsed dict) for the CLI to compose with `read_config_file()` and + `from_dict()`. - `mace_core.metadata` — `ModelMetadata`, the versioned record every trained model carries (config as written and as resolved, provenance, a summary per data source, per head its E0s and the sources it consumed, DOI, citations, diff --git a/packages/mace-core/src/mace_core/config/__init__.py b/packages/mace-core/src/mace_core/config/__init__.py index 9a507e7b1..c9178ccfa 100644 --- a/packages/mace-core/src/mace_core/config/__init__.py +++ b/packages/mace-core/src/mace_core/config/__init__.py @@ -1,21 +1,23 @@ """Configuration schemas for MACE v1. -The base machinery lives in `base`; the training schema sections (model, data, -training, ...) arrive with their own tickets and are re-exported from here. +The base machinery lives in `base` and the command-line override grammar in +`cli`; the training schema sections (model, data, training, ...) arrive with +their own tickets and are re-exported from here. """ from mace_core.config.base import ( ConfigError, ConfigSection, - ConfigWarning, ReforgeBaseConfig, read_config_file, ) +from mace_core.config.cli import apply_overrides, parse_overrides __all__ = [ "ConfigError", "ConfigSection", - "ConfigWarning", "ReforgeBaseConfig", + "apply_overrides", + "parse_overrides", "read_config_file", ] diff --git a/packages/mace-core/src/mace_core/config/_schema_rules.py b/packages/mace-core/src/mace_core/config/_schema_rules.py deleted file mode 100644 index 0358c731a..000000000 --- a/packages/mace-core/src/mace_core/config/_schema_rules.py +++ /dev/null @@ -1,240 +0,0 @@ -"""The schema rules a config class must keep, checked when the class is defined, -and the annotation introspection the loader shares with them. - -A *section* is a `ConfigSection` subclass; a *kinds field* holds a union of -two or more sections, declared `Annotated[A | B, Field(discriminator="kind")]` -with `kind: Literal["a"]` in each variant. Every rule raises `TypeError` -`. ` naming the fix. - -`base` is bound as a module and dereferenced at call time only: `base.py` -imports this module at its top, so a `from mace_core.config.base import ...` -here would break the package import whichever module is imported first. -""" - -from __future__ import annotations - -import types -from typing import Annotated, Any, Literal, TypeGuard, Union, get_args, get_origin - -from pydantic import BaseModel -from pydantic_core import PydanticUndefined - -from mace_core.config import base as _base - -_NO_KIND_FIX = ( - "default to a variant, or for none of the kinds declare an empty variant " - "with kind: Literal['none']" -) - - -def is_section(node: Any) -> TypeGuard[type[_base.ConfigSection]]: - """A section class. A parametrised generic such as `list[int]` passes - `isinstance(node, type)` on 3.10, hence the origin check first.""" - return ( - get_origin(node) is None - and isinstance(node, type) - and issubclass(node, _base.ConfigSection) - ) - - -def unwrap(node: Any) -> Any: - """`Annotated[X, ...]` as `X`; anything else unchanged.""" - return get_args(node)[0] if get_origin(node) is Annotated else node - - -def members_of(node: Any) -> tuple[Any, ...]: - """The members of a union, nested unions and `Annotated` flattened; a - non-union is its own single member.""" - node = unwrap(node) - if get_origin(node) in (Union, types.UnionType): - return tuple(member for arg in get_args(node) for member in members_of(arg)) - return (node,) - - -def tag_of(variant: type[BaseModel]) -> str | None: - """The one string of a variant's `kind: Literal[...]`; None for any other - shape of tag.""" - field = variant.model_fields.get("kind") - if field is None or get_origin(field.annotation) is not Literal: - return None - values = get_args(field.annotation) - return values[0] if len(values) == 1 and isinstance(values[0], str) else None - - -def kinds_of(node: Any) -> dict[str, type[_base.ConfigSection]] | None: - """`{tag: variant}` for a union of two or more sections, in declaration - order; None for anything else. A variant without a proper tag is keyed by - its class name so that `check_kinds_field` can still name it.""" - members = members_of(node) - if len(members) < 2 or not all(is_section(member) for member in members): - return None - return {tag_of(member) or member.__name__: member for member in members} - - -def kinds_fields_of(cls: type[BaseModel]) -> dict[str, dict[str, type[Any]]]: - """The kinds fields of a class, each with its `{tag: variant}`.""" - fields = {} - for name, field in cls.model_fields.items(): - kinds = kinds_of(field.annotation) - if kinds is not None: - fields[name] = kinds - return fields - - -def reachable_sections(cls: type[BaseModel]) -> list[type[BaseModel]]: - """Every section reachable from `cls` through its fields, `cls` included. - A class left incomplete by a forward reference is rebuilt on the way, which - raises if the name never resolves (F7).""" - found: list[type[BaseModel]] = [] - pending = [cls] - while pending: - section = pending.pop() - if section in found: - continue - if not section.__pydantic_complete__: - section.model_rebuild() - found.append(section) - for field in section.model_fields.values(): - pending.extend(sections_in(field.annotation)) - return found - - -def sections_in(annotation: Any) -> list[type[BaseModel]]: - """The sections inside an annotation: union members and the parameters - of lists, dicts and tuples, at any depth.""" - node = unwrap(annotation) - if is_section(node): - return [node] - return [section for arg in get_args(node) for section in sections_in(arg)] - - -def admits_none(annotation: Any) -> bool: - return any(m in (Any, object, type(None)) for m in members_of(annotation)) - - -def check_model_config(cls: type[BaseModel]) -> None: - """G5: `extra="forbid"` cannot be reopened by a subclass.""" - extra = cls.model_config.get("extra") - if extra != "forbid": - raise TypeError( - f"{cls.__name__} sets extra={extra!r}; a section keeps extra='forbid' " - f"so that an unknown key is an error" - ) - - -def check_fields(cls: type[BaseModel]) -> None: - """F1-F6 over every field of a complete class.""" - for name in cls.model_computed_fields: - raise TypeError( - f"{cls.__name__}.{name} is a computed field; the resolved export " - f"could not be loaded back" - ) - for name, field in cls.model_fields.items(): - where = f"{cls.__name__}.{name}" - if field.alias or field.validation_alias or field.serialization_alias: - raise TypeError( - f"{where} has an alias; the resolved export could not be loaded back" - ) - if field.exclude: - raise TypeError( - f"{where} is excluded from dumps; the resolved export could not be " - f"loaded back" - ) - check_annotation(where, field.annotation, field.discriminator is not None) - kinds = kinds_of(field.annotation) - if kinds is not None: - check_kinds_field(where, kinds, field.default) - elif field.default is None and not admits_none(field.annotation): - raise TypeError(f"{where} defaults to None, which its type does not admit") - - -def _is_lenient_model(node: Any) -> bool: - return ( - get_origin(node) is None - and isinstance(node, type) - and issubclass(node, BaseModel) - and not is_section(node) - ) - - -def check_annotation( - where: str, annotation: Any, discriminated: bool, inside: bool = False -) -> None: - """F1 (no sets), F4 (sections only) and F5 (union shapes) at every depth; - `inside` is true under a dict, list or tuple, where a kinds union is not - allowed.""" - node = unwrap(annotation) - if get_origin(node) is Literal: - return - if node in (set, frozenset) or get_origin(node) in (set, frozenset): - raise TypeError( - f"{where} is typed as a set, whose order changes between runs. Use a list" - ) - members = members_of(node) - for member in members: - if _is_lenient_model(member): - raise TypeError( - f"{where} holds {member.__name__}, which is not a ConfigSection" - ) - sections = [member for member in members if is_section(member)] - if sections: - if type(None) in members: - raise TypeError(f"{where} admits None, which is not a kind; {_NO_KIND_FIX}") - others = [member for member in members if not is_section(member)] - if others: - other = get_origin(others[0]) or others[0] - raise TypeError( - f"{where} mixes its kinds with {other.__name__}; a kinds field holds " - f"its variants only" - ) - if len(sections) > 1 and inside: - raise TypeError( - f"{where} is a union of sections inside a dict, list or tuple; " - f"declare a section holding the kinds field there" - ) - if len(sections) > 1 and not discriminated: - raise TypeError( - f"{where} is a union of sections without a discriminator; declare " - f"kinds as Annotated[A | B, Field(discriminator='kind')]" - ) - return - if len(members) > 1: - for member in members: - check_annotation(where, member, False, inside) - else: - for arg in get_args(node): - check_annotation(where, arg, False, True) - - -def check_kinds_field(where: str, kinds: dict[str, type[Any]], default: Any) -> None: - """F6: one-string tags, no variant named like the tag, no variant field - named like a kind (the flat form could not tell it from a kind's settings), - a variant as the default. A default factory is neither run nor checked here - (F3); its result is checked when an empty mapping selects it.""" - for tag, variant in kinds.items(): - if tag_of(variant) is None: - raise TypeError( - f"{where} has variant {variant.__name__} whose kind must be a " - f"Literal of exactly one string" - ) - if tag == "kind": - raise TypeError( - f"{where} has variant {variant.__name__} whose kind is named 'kind' " - f"like the tag; rename it" - ) - for name in variant.model_fields: - if name != "kind" and name in kinds: - raise TypeError( - f"{where} has variant {variant.__name__} with a field named " - f"like the kind {name}; rename the field" - ) - if default is None: - raise TypeError( - f"{where} defaults to None, which is not a kind; {_NO_KIND_FIX}" - ) - if default is not PydanticUndefined and type(default) not in kinds.values(): - example = next(iter(kinds.values())).__name__ - raise TypeError( - f"{where} has a default that is not one of its variants; write e.g. " - f"{example}()" - ) diff --git a/packages/mace-core/src/mace_core/config/base.py b/packages/mace-core/src/mace_core/config/base.py index 514a95455..534662d32 100644 --- a/packages/mace-core/src/mace_core/config/base.py +++ b/packages/mace-core/src/mace_core/config/base.py @@ -1,621 +1,199 @@ -"""The config base: files and dotted command-line overrides into one validated -pydantic tree, and the tree back out as a dict. - -`load` turns every file (any number, in order) and every override into leaf -updates `(path, value, source)`, walks each path once through the schema (the -one schema-dependent step: unknown keys and wrong shapes are rejected with full -dotted paths, a bare kind name at a kinds field becomes its `kind`, and the -kinds fields the path enters are recorded on the update), merges them -set-at-path into one plain dict, validates that dict once, and finally warns -for each override that lost its effect. Precedence is defaults < files in -order < overrides in order. Nothing else feeds a config: no environment, no -dotenv. - -Wire form of a kinds field (`Annotated[A | B, Field(discriminator="kind")]`): -a mapping with a kind name holding that kind's settings (`huber: {delta: 0.1}`, -any number of them) and `kind: huber` selecting the one that runs; a bare -`huber` is `{kind: huber}`; a single kind key selects itself; `{}` keeps the -schema default. Two or more kind keys without a selection, a required kinds -field with nothing written, and the tag under a kind are the one kinds error, -rendered with the path, the fix and the source of each kind key. The exports -write the kind that ran as `{huber: {fields}}`, without the tag. `kind` is the -tag only under a kinds field: the same class as a plain section field keeps -`kind` as an ordinary key on input and export. - -An override loses its effect in two ways, and only these warn (a file never -does): a later override writes at its path, above it, or below a value of its -that was not a mapping (`{}` never clears a mapping); or it wrote under a kind -that does not run at its kinds field. +"""One config file, one validation, two exports. + +`ReforgeBaseConfig.load(config_file)` reads one TOML, YAML or JSON file into a +dict (`read_config_file`, public) and validates it once with pydantic +(`from_dict`, public, for a caller that edits the dict first: `cli` does, for a +command line). Everything the schema rejects, unknown keys included, is +pydantic's `ValidationError` at its dotted location; `ConfigError` is raised +only for a file that cannot be read or parsed. `to_resolved_dict` is the full +JSON-native dump, itself a reloadable config file; `to_user_dict` holds only +what was set. A kinds field (a discriminated union on `kind`) is an ordinary +pydantic feature written in the tagged form, `loss: {kind: huber, delta: 0.1}`; +nothing here knows about it. Nothing here knows about a command line either; +the `--a.b value` grammar is `cli`'s. """ -from __future__ import annotations - -import copy -import difflib import json +import os import sys -import warnings -from collections.abc import Collection, Iterable +from collections.abc import Iterator, Mapping +from collections.abc import Set as AbstractSet from pathlib import Path -from typing import Any, NamedTuple, TypeVar, get_args, get_origin +from typing import Any, get_args, get_origin import yaml -from pydantic import ( - BaseModel, - ConfigDict, - SerializerFunctionWrapHandler, - ValidationError, - model_serializer, - model_validator, -) -from pydantic_core import ( - ErrorDetails, - PydanticCustomError, - PydanticUndefined, -) +from pydantic import BaseModel, ConfigDict from typing_extensions import Self -from mace_core.config._schema_rules import ( - check_fields, - check_model_config, - is_section, - kinds_fields_of, - kinds_of, - members_of, - reachable_sections, - unwrap, -) - if sys.version_info >= (3, 11): import tomllib else: import tomli as tomllib -_PARSERS = { - ".toml": tomllib.loads, - ".yaml": yaml.safe_load, - ".yml": yaml.safe_load, - ".json": json.loads, -} -#: The one error `_to_tagged_form` raises through pydantic; `load` renders it. -_KINDS_ERROR = "kinds" +__all__ = ["ConfigError", "ConfigSection", "ReforgeBaseConfig", "read_config_file"] class ConfigError(ValueError): - """A file or command line the schema cannot take; the message names the - key and the fix.""" - - -class ConfigWarning(UserWarning): - """An override that had no effect on the config that runs.""" - - -class KindEntered(NamedTuple): - """A kinds field an update's path passes through under one of its kind - keys: the field's path, its kind names and the key entered.""" - - field: tuple[Any, ...] - kinds: tuple[str, ...] - key: str - - -class Update(NamedTuple): - """One leaf update; `source` is the file's path or the override as typed. - The walk fills `under_kinds` with every kinds field the path enters under a - kind key, and `selects` with the kinds field whose tag the path ends at.""" - - path: tuple[Any, ...] - value: Any - source: str - under_kinds: tuple[KindEntered, ...] = () - selects: tuple[Any, ...] | None = None - - -#: The node for the `kind` slot under a kinds field. -_TAG = object() - - -class _KindPosition(NamedTuple): - """A variant class reached through its kinds field, where `kind` is the - tag and not a key. The same class as a plain field is walked as a section.""" - - variant: type[ConfigSection] - - -def _narrowed(node: Any) -> Any: - """`Annotated` stripped and `X | None` stepped through to `X`.""" - members = [m for m in members_of(node) if m is not type(None)] - return unwrap(members[0]) if len(members) == 1 else unwrap(node) - - -def _step(node: Any, key: Any) -> tuple[Any, list[Any] | None]: - """One step of the walk: the child node under `key` (None when the key - is not valid there) and the keys valid at `node` (None when any key is, - `[]` when the node is written whole).""" - node = unwrap(node) - if isinstance(node, _KindPosition): - fields = node.variant.model_fields - keys = [name for name in fields if name != "kind"] - return (fields[key].annotation if key in keys else None), keys - if (kinds := kinds_of(node)) is not None: - keys = ["kind", *kinds] - if key == "kind": - return _TAG, keys - return (_KindPosition(kinds[key]) if key in kinds else None), keys - if is_section(node): - fields = node.model_fields - return (fields[key].annotation if key in fields else None), list(fields) - node = _narrowed(node) - origin = get_origin(node) - if origin is dict: - return get_args(node)[1], None - if node in (Any, object, dict): - return Any, None - if origin in (list, tuple) or node in (list, tuple): - if not isinstance(key, int): - return None, [] - args = get_args(node) - if origin is tuple and Ellipsis not in args and args: - return (args[key] if key < len(args) else Any), [] - return (args[0] if args else Any), [] - return None, [] - - -def _dotted(path: tuple[Any, ...]) -> str: - return ".".join(map(str, path)) - - -def _as_shown(value: Any) -> str: - """JSON where it can, `str` where it cannot (a TOML date, a YAML set).""" - return json.dumps(value, default=str) - - -def _as_typed(value: Any) -> str: - """A command-line value as typed (a str), anything else as JSON.""" - return value if isinstance(value, str) else _as_shown(value) - - -def _unknown_key( - parent: Any, keys: list[Any] | None, path: tuple[Any, ...] -) -> ConfigError: - """D1: the key by its dotted path, then the nearest neighbour, the kinds - of a kinds field, or why the tag is not a key under a kind.""" - message = f"unknown config key '{_dotted(path)}'" - above = _dotted(path[:-1]) - if keys == []: - return ConfigError( - f"{message}; {above} is written whole and takes no keys under it" - ) - nearest = difflib.get_close_matches( - str(path[-1]), [str(k) for k in keys or []], n=1 - ) - if nearest: - message += f"; did you mean '{_dotted((*path[:-1], nearest[0]))}'?" - parent = unwrap(parent) - if (kinds := kinds_of(parent)) is not None: - message += f"; the kinds of {above} are {', '.join(kinds)}" - elif isinstance(parent, _KindPosition) and path[-1] == "kind": - message += f"; the key {path[-2]} already names the kind" - return ConfigError(message) - - -def _flatten(value: Any, path: tuple[Any, ...], source: str) -> list[Update]: - """Leaf updates of a parsed document: a non-empty mapping recurses, anything - else (a scalar, a whole list, an empty mapping) is a leaf. Each leaf is a - copy, so a YAML anchor shared under two keys becomes two values (G4).""" - if isinstance(value, dict) and value: - return [ - u - for key, item in value.items() - for u in _flatten(item, (*path, key), source) - ] - return [Update(path, copy.deepcopy(value), source)] - - -def _resolve( - root: type[ConfigSection], update: Update, problems: list[str] -) -> Update | None: - """Walk one update's path through the schema. A bare kind name at a kinds - field comes back as its `kind` update; a list value is checked item by item - but stays whole. A `ConfigError` is recorded under the update's source and - the update dropped, so that every problem of one load is reported together.""" - path, value, source = update.path, update.value, update.source - under_kinds: list[KindEntered] = [] - try: - parent, node = None, root - for depth, key in enumerate(path): - child, keys = _step(node, key) - if child is None: - raise _unknown_key(node, keys, path[: depth + 1]) - if isinstance(child, _KindPosition): - kinds = tuple(kinds_of(unwrap(node)) or ()) - under_kinds.append(KindEntered(path[:depth], kinds, key)) - parent, node = node, child - if node is _TAG: - kinds = kinds_of(unwrap(parent)) or {} - if not isinstance(value, str): - raise ConfigError( - f"{_dotted(path)} must name a kind, one of {', '.join(kinds)}; " - f"got {_as_shown(value)}" - ) - path, node = path[:-1], parent - node = unwrap(node) - if (kinds := kinds_of(node)) is not None: - if isinstance(value, str): - if value not in kinds: - raise _unknown_key(node, ["kind", *kinds], (*path, value)) - return Update((*path, "kind"), value, source, tuple(under_kinds), path) - if value != {}: - raise ConfigError( - f"{_dotted(path)} must be the name of a kind or a mapping under " - f"one, one of {', '.join(kinds)}; got {_as_shown(value)}" - ) - elif isinstance(node, _KindPosition) or is_section(node): - if not isinstance(value, dict): - raise ConfigError( - f"{_dotted(path)} must be a mapping of its keys; " - f"got {_as_shown(value)}" - ) - else: - node = _narrowed(node) - is_list = get_origin(node) in (list, tuple) or node in (list, tuple) - if is_list and isinstance(value, list): - for index, item in enumerate(value): - for item_update in _flatten(item, (*path, index), source): - _resolve(root, item_update, problems) - return Update(path, value, source, tuple(under_kinds)) - except ConfigError as error: - problems.append(f"{source}: {error}") - return None - - -def _resolve_all(root: type[ConfigSection], updates: list[Update]) -> list[Update]: - problems: list[str] = [] - resolved = [_resolve(root, update, problems) for update in updates] - if problems: - raise ConfigError("\n".join(problems)) - return [update for update in resolved if update is not None] - - -def _set_at_path(mapping: dict[Any, Any], path: tuple[Any, ...], value: Any) -> None: - """Set `value` at `path`, creating mappings on the way; an empty mapping - never clears a mapping already there, and is stored as a fresh one so that - later writes under it leave the update's value alone.""" - if not path: - return - for key in path[:-1]: - if not isinstance(mapping.get(key), dict): - mapping[key] = {} - mapping = mapping[key] - if value != {} or not isinstance(mapping.get(path[-1]), dict): - mapping[path[-1]] = {} if value == {} else value - - -def read_config_file(config_file: str | Path) -> dict[str, Any]: - """Parse one file by its extension (`.toml`, `.yaml`, `.yml`, `.json`, in - any case). - - A `ConfigError` names the file for an unknown extension, a file that cannot - be read or parsed, and a top level that is not a table. An empty or - comment-only file is `{}`. Non-string YAML keys (`1:`) stay as parsed; - quote them where the schema wants strings. - """ - path = Path(config_file) - parser = _PARSERS.get(path.suffix.lower()) - if parser is None: + """A config file, or a command-line override (`cli`), that cannot be read or + parsed. Schema errors are pydantic's.""" + + +def read_config_file(path: str | os.PathLike[str]) -> dict[str, Any]: + """Parse one TOML, YAML or JSON config file, chosen by its (case-folded) + extension, into a dict; an empty or comment-only file is `{}`. Raises + `ConfigError` for an unknown extension, an unreadable file, a parse failure, + a top level that is not a mapping, or a value that contains itself (a YAML + anchor inside itself), which no schema could export again.""" + path = Path(path) + parsers = { + ".toml": tomllib.loads, + ".yaml": yaml.safe_load, + ".yml": yaml.safe_load, + ".json": json.loads, + } + parse = parsers.get(path.suffix.lower()) + if parse is None: raise ConfigError( - f"unknown extension '{path.suffix}' of config file {path}; use .toml, " - f".yaml, .yml or .json" + f"cannot read config file {path}: unknown extension {path.suffix!r}; " + "use .toml, .yaml, .yml or .json" ) try: text = path.read_text(encoding="utf-8") except (OSError, UnicodeDecodeError) as error: raise ConfigError(f"cannot read config file {path}: {error}") from error try: - document = parser(text) - except (ValueError, yaml.YAMLError) as error: + document = parse(text) + except (ValueError, yaml.YAMLError, RecursionError) as error: # toml, json: Value raise ConfigError(f"cannot parse config file {path}: {error}") from error - if document is None: - document = {} + if document is None: # YAML reads an empty or comment-only file as None + return {} if not isinstance(document, dict): raise ConfigError( - f"config file {path} must have a table of keys at the top level, not " - f"{type(document).__name__}" + f"config file {path} must be a mapping of keys to values at the top " + f"level, not {type(document).__name__}" ) + try: # the stdlib's cycle detector; shared siblings pass, only a cycle fails + json.dumps(document, default=str, skipkeys=True) + except ValueError as error: + raise ConfigError( + f"config file {path} contains a value that refers to itself" + ) from error return document -def _parse_overrides(tokens: Iterable[str]) -> list[Update]: - """`--a.b.c value` or `--a.b.c=value`, in order. A value that is `null` or - starts with `[` or `{` is JSON (and flattened like a file); any other value - stays a string for pydantic to coerce. Keys match exactly.""" - argv = list(tokens) - updates = [] - position = 0 - while position < len(argv): - token = argv[position] - position += 1 - key, has_inline_value, value = token[2:].partition("=") - if not token.startswith("--") or not key: - raise ConfigError( - f"unknown config option '{token}'; options are --key.path value" - ) - if not has_inline_value: - if position == len(argv): - raise ConfigError(f"override {token} is missing its value") - value = argv[position] - position += 1 - source = token if has_inline_value else f"{token} {value}" - parsed: Any = value - if value == "null" or value[:1] in ("[", "{"): - try: - parsed = json.loads(value) - except ValueError as error: - raise ConfigError( - f"override {token} is not valid JSON: {error}" - ) from error - updates.extend(_flatten(parsed, tuple(key.split(".")), source)) - return updates - - class ConfigSection(BaseModel): - """A node of a config tree: unknown keys are errors, and the schema rules - of `_schema_rules` are checked when a subclass is defined (or, behind a - forward reference, on the first `load`). - - A kinds field takes the wire form `{kind: X, X: {...}, Y: {...}}`, `{X: - {...}}`, `"X"` or `{}` and dumps as `{X: {fields}}`; internally pydantic - sees the tagged form `{kind: X, ...X's fields}`, which `model_validate` and - direct construction accept as well. An instance of a subclass of a variant - is held as given and exported under the variant's kind. + """A node of a config tree: unknown keys are errors at every level. + + Subclasses declare fields only. The definition-time checks keep unknown + keys fatal (every reachable model is a section, `extra="forbid"` cannot be + reopened) and the resolved export a fixed point (no sets, aliases, + excluded or computed fields). A class left incomplete by a forward + reference is checked on every `load` of a root that reaches it instead; + a forward reference to a class local to a function cannot be resolved + that way (pydantic cannot see the function's scope), so declare such + classes at module level. """ - # inf/nan are exported as floats (JSON constants Infinity, NaN) rather than - # pydantic's default null, which would store a different value in the model - # metadata; `metadata._Record` writes them the same way. + # inf/nan as the JSON constants Infinity and NaN, not pydantic's default + # null, which would turn a value into a different one. model_config = ConfigDict(extra="forbid", ser_json_inf_nan="constants") @classmethod def __pydantic_init_subclass__(cls, **kwargs: Any) -> None: super().__pydantic_init_subclass__(**kwargs) - check_model_config(cls) - if cls.__pydantic_complete__: - check_fields(cls) - - @model_validator(mode="before") - @classmethod - def _kinds_fields_to_tagged_form(cls, data: Any) -> Any: - if not isinstance(data, dict): - return data - for name, kinds in kinds_fields_of(cls).items(): - if name in data: - data = {**data, name: _to_tagged_form(cls, name, data[name], kinds)} - return data - - # No return annotation: with one, the serialization JSON schema collapses. - @model_serializer(mode="wrap") - def _kinds_fields_under_their_name(self, handler: SerializerFunctionWrapHandler): - dumped = handler(self) - for name, kinds in kinds_fields_of(type(self)).items(): - if name in dumped: # absent under exclude_unset - variant = getattr(self, name) - if not isinstance(variant, tuple(kinds.values())): - raise TypeError( - f"{type(self).__name__}.{name} holds {type(variant).__name__}, " - f"which is not one of its variants" - ) - settings = {k: v for k, v in dumped[name].items() if k != "kind"} - dumped[name] = {variant.kind: settings} - return dumped - - -def _selected_kind(value: dict[str, Any], kinds: Collection[str]) -> str | None: - """The kind a kinds field's wire mapping runs: its `kind` scalar, else its - sole kind key; None when nothing selects one (an empty mapping runs the - schema default, several kind keys need a selection).""" - selected = value.get("kind") - if isinstance(selected, str): - return selected - written = [key for key in value if key in kinds] - return written[0] if len(written) == 1 else None + extra = cls.model_config.get("extra") + if extra != "forbid": + raise TypeError( + f"{cls.__name__} sets extra={extra!r}; a section keeps " + "extra='forbid' so an unknown key stays an error" + ) + if cls.__pydantic_complete__: # else a forward reference: `load` checks it + _check_field_declarations(cls) -def _kinds_error( - field: str, kinds: list[str], under: str | None = None -) -> PydanticCustomError: - """The kinds error: `field` needs a kind (one of `kinds`), or the tag was - written under the kind `under`. `load` renders it with the path, the fix - and the sources (`_kind_error_message`).""" - choices = " or ".join(f"kind: {kind}" for kind in kinds) - template = ( - "{field} needs a kind; write {choices}" - if under is None - else "{field}.{under}.kind is not a key; {under} already names the kind" - ) - context = {"field": field, "kinds": kinds, "under": under, "choices": choices} - return PydanticCustomError(_KINDS_ERROR, template, context) +def _leaf_types(annotation: Any) -> Iterator[Any]: + """Every class or origin an annotation reaches through `Annotated`, unions and + list, tuple or dict parameters. Non-type arguments (Literal values, `Field` + metadata) come out too; callers test `isinstance(leaf, type)`.""" + yield get_origin(annotation) or annotation + for argument in get_args(annotation): + yield from _leaf_types(argument) -def _to_tagged_form( - cls: type[ConfigSection], name: str, value: Any, kinds: dict[str, type[Any]] -) -> Any: - """Wire form of one kinds field to pydantic's tagged form. Shapes the - schema cannot take are returned unchanged for pydantic to report; the - shapes it would misreport raise the kinds error, among them a `kind` inside - the selected kind's settings, which is not a key there.""" - if isinstance(value, str): - return {"kind": value} - if not isinstance(value, dict) or not isinstance(value.get("kind", ""), str): - return value - selected = _selected_kind(value, kinds) - if selected is None: - written = [key for key in value if key in kinds] - if written: - raise _kinds_error(name, written) - if value: - return value - default = cls.model_fields[name].get_default(call_default_factory=True) - if default is PydanticUndefined: - raise _kinds_error(name, list(kinds)) - if type(default) not in kinds.values(): - example = next(iter(kinds.values())).__name__ - raise TypeError( - f"{cls.__name__}.{name} has a default factory whose result is not " - f"one of its variants; return e.g. {example}()" - ) - return {"kind": default.kind, **default.model_dump(exclude_unset=True)} - settings = value.get(selected, {}) - if not isinstance(settings, dict): - return value - if "kind" in settings: - raise _kinds_error(name, list(kinds), under=selected) - rest = {k: v for k, v in value.items() if k not in kinds and k != "kind"} - return {**rest, **settings, "kind": selected} +def _check_field_declarations(section: type[ConfigSection]) -> None: + """The checks on one class's own fields; a violation is a `TypeError`.""" + if section.model_computed_fields: + name = next(iter(section.model_computed_fields)) + raise TypeError( + f"{section.__name__}.{name} is a computed field; it would not load back" + ) + for name, field in section.model_fields.items(): + where = f"{section.__name__}.{name}" + if field.alias or field.validation_alias or field.serialization_alias: + raise TypeError(f"{where} has an alias; a config key is its field name") + if field.exclude: + raise TypeError(f"{where} is excluded from dumps; it would not load back") + for leaf in _leaf_types(field.annotation): + if not isinstance(leaf, type): + continue + # Any set type, `set`, `frozenset` or an abstract one, dumps in an + # order that varies between runs, so the export would not be stable. + if issubclass(leaf, AbstractSet): + raise TypeError(f"{where} is typed as a set; order varies. Use a list") + if issubclass(leaf, BaseModel) and not issubclass(leaf, ConfigSection): + raise TypeError( + f"{where} holds {leaf.__name__}, which is not a ConfigSection; " + "unknown keys under it would be dropped" + ) -_ConfigT = TypeVar("_ConfigT", bound="ReforgeBaseConfig") +def _check_sections_reached_by(root: type[ConfigSection]) -> None: + """Resolve any section the root's tree reaches that a forward reference left + incomplete at definition, and check every reached section. Checking again on + each load is a few attribute reads per class; it saves remembering which + classes were checked.""" + to_visit, seen = [root], set() + while to_visit: + section = to_visit.pop() + if section in seen: + continue + seen.add(section) + if not section.__pydantic_complete__: + section.model_rebuild() # resolves the forward reference or raises + _check_field_declarations(section) + for field in section.model_fields.values(): + for leaf in _leaf_types(field.annotation): + if isinstance(leaf, type) and issubclass(leaf, ConfigSection): + to_visit.append(leaf) class ReforgeBaseConfig(ConfigSection): - """The root of a config tree: `load` builds it from files and the command - line; the two exports write it back as JSON-native dicts.""" + """The root of a config tree: `load` a file, or `from_dict` a parsed one.""" @classmethod - def load( - cls, - config_files: str | Path | Iterable[str | Path] | None = (), - cli_overrides: Iterable[str] = (), - ) -> Self: - """Build the config from `config_files` in order, then `cli_overrides` - in order (`sys.argv[1:]`-style tokens); a lone path is one file, None - is no file. + def load(cls, config_file: str | os.PathLike[str]) -> Self: + """Build the config from one TOML, YAML or JSON file.""" + return cls.from_dict(read_config_file(config_file)) - Raises `ConfigError` for a file or override the schema cannot take - (every unknown key of the load together, each under its file or - override), pydantic's `ValidationError` for a value of the wrong type, - and warns `ConfigWarning` for each override that had no effect on the - config that runs. - """ - if isinstance(cli_overrides, str): - raise TypeError( - "cli_overrides is a string; pass the tokens as a list, " - "like sys.argv[1:]" - ) - for section in reachable_sections(cls): - check_fields(section) - if config_files is None: - config_files = () - files = ( - [config_files] - if isinstance(config_files, (str, Path)) - else list(config_files) - ) - from_files = [ - update - for file in files - for update in _flatten(read_config_file(file), (), str(Path(file))) - ] - updates = _resolve_all(cls, [*from_files, *_parse_overrides(cli_overrides)]) - merged: dict[str, Any] = {} - for update in updates: - _set_at_path(merged, update.path, update.value) - config = _validated(cls, merged, updates) - overrides = updates[len(from_files) :] - messages = [ - _override_without_effect(update, overrides[position + 1 :], merged) - for position, update in enumerate(overrides) - ] - for message in dict.fromkeys(m for m in messages if m is not None): - warnings.warn(message, ConfigWarning, stacklevel=2) - return config + @classmethod + def from_dict(cls, document: Mapping[str, Any]) -> Self: + """Build the config from a parsed document, checking the sections the + class reaches first. This is `load` for a caller that edits the parsed + file before validating it, such as a command line writing its flags into + the dict `read_config_file` returned. The document is not written into + (values under an `Any`-typed field are shared with it, not copied).""" + _check_sections_reached_by(cls) + return cls.model_validate(document) def to_resolved_dict(self) -> dict[str, Any]: - """Every field with defaults filled, JSON-native, in declaration order; - loading it back gives an equal config and the same dict. A kinds field - holding a subclass of its variant is written as the variant, so a - field the subclass added is not written.""" + """Every field, defaults filled, as JSON-native values in declaration + order: a config file that loads back to this config (through TOML + whenever no value is `None`).""" return self.model_dump(mode="json") def to_user_dict(self) -> dict[str, Any]: - """Only what the files and the overrides set, in the same shape; the - kind that ran is recorded even when it was chosen by default.""" + """Only what the file set, in the same shape. A loaded config carries + the tag of every kinds field it wrote, so this loads back; + a variant built in code without its tag (`Config(loss=Huber(delta=2))`) + exports without `kind`: pass the tag, or use `to_resolved_dict`.""" return self.model_dump(mode="json", exclude_unset=True) - - -def _validated( - cls: type[_ConfigT], merged: dict[str, Any], updates: list[Update] -) -> _ConfigT: - """Validate once; a kinds error comes out as one `ConfigError`, and a load - without a kinds error passes pydantic's `ValidationError` through (D3). A - kinds error ends the validation of its class, so the other errors of that - class are reported on the next load.""" - try: - return cls.model_validate(merged) - except ValidationError as error: - kind_errors = [e for e in error.errors() if e["type"] == _KINDS_ERROR] - if not kind_errors: - raise - lines = [_kind_error_message(e, updates) for e in kind_errors] - raise ConfigError("\n".join(lines)) from error - - -def _wrote(update: Update, target: tuple[Any, ...]) -> bool: - """Whether the update wrote at or under `target`, or a whole list holding it.""" - depth = len(update.path) - if depth >= len(target): - return update.path[: len(target)] == target - return update.path == target[:depth] and isinstance(update.value, list) - - -def _kind_error_message(error: ErrorDetails, updates: list[Update]) -> str: - """Pydantic's location is the section that raised, and the field is in the - context. The dotted fix is left out inside a list item, where the walk - would reject it; each kind that was written is named with the source that - first wrote it. The tag under a kind is rejected by the walk before - validation; should it arrive, pydantic's own sentence is kept.""" - ctx = error.get("ctx") or {} - if ctx.get("under") is not None: - return error["msg"] - path = (*error["loc"], ctx["field"]) - kinds = ctx["kinds"] - message = f"{_dotted(path)} needs a kind; write {ctx['choices']} in a file" - if not any(isinstance(part, int) for part in path): - message += f", or pass --{_dotted(path)} {kinds[0]}" - writers = [] - for kind in kinds: - writer = next((u for u in updates if _wrote(u, (*path, kind))), None) - if writer is not None: - writers.append(f"{kind} from {writer.source}") - if writers: - message += f"; {', '.join(writers)}" - return message - - -def _override_without_effect( - update: Update, later: list[Update], merged: dict[str, Any] -) -> str | None: - """C11, the two ways an override loses its effect: a later override writes - at its path, above it, or below a value of its that was not a mapping (`{}` - never clears a mapping); or it wrote under a kind that does not run at its - kinds field, decided from the merged dict as `_to_tagged_form` decides it. - Never raises.""" - for other in reversed(later): - depth = min(len(update.path), len(other.path)) - if other.value == {} or update.path[:depth] != other.path[:depth]: - continue - if len(other.path) > len(update.path) and update.value == {}: - continue - if other.selects is not None: - field, running = _dotted(other.selects), _as_typed(other.value) - return f"{update.source} is overridden: {field} runs {running}" - at, value = _dotted(other.path), _as_typed(other.value) - return f"{update.source} is overridden: {at} is {value}" - for field, kinds, key in update.under_kinds: - mapping: Any = merged - for step in field: - mapping = mapping.get(step) if isinstance(mapping, dict) else None - running = _selected_kind(mapping, kinds) if isinstance(mapping, dict) else None - if running is not None and running != key: - at = _dotted(field) - return f"{update.source}: {at}.{key} has no effect, {at} runs {running}" - return None diff --git a/packages/mace-core/src/mace_core/config/cli.py b/packages/mace-core/src/mace_core/config/cli.py new file mode 100644 index 000000000..48c9becd8 --- /dev/null +++ b/packages/mace-core/src/mace_core/config/cli.py @@ -0,0 +1,108 @@ +"""The command-line half of a config: `--a.b value` tokens, and a parsed file +edited before validation. + +`parse_overrides(tokens)` turns command-line tokens into a mapping of dotted +path to value; `apply_overrides(document, overrides)` writes such a mapping +into a copy of a parsed config file. A command line then builds its config as + + Config.from_dict(apply_overrides(read_config_file(path), parse_overrides(rest))) + +and every error the schema reports is pydantic's, at the dotted location the +override named. Nothing in `base` knows about this module; a CLI that exposes +explicit flags instead can hand `apply_overrides` a mapping it built itself. +""" + +import json +from collections.abc import Iterable, Mapping +from typing import Any + +from mace_core.config.base import ConfigError + +__all__ = ["apply_overrides", "parse_overrides"] + + +def parse_overrides(tokens: Iterable[str]) -> dict[str, Any]: + """`--a.b.c value` or `--a.b.c=value`, in order, to `{"a.b.c": value}`. A + value that is `null` or starts with `[` or `{` is parsed as JSON; any other + value stays a string for pydantic to coerce. A repeated path keeps its last + value, at its last position. A token that is not `--path`, a path without + its value or with an empty key (`--a..b`), and a JSON value that does not + parse are `ConfigError`s.""" + if isinstance(tokens, str): + raise TypeError(f"tokens is a string, {tokens!r}; pass a list of tokens") + argv = list(tokens) + overrides: dict[str, Any] = {} + position = 0 + while position < len(argv): + token = argv[position] + position += 1 + dotted_path, has_inline_value, value = token[2:].partition("=") + if not token.startswith("--") or not dotted_path: + raise ConfigError( + f"unknown config option '{token}'; options are --key.path value" + ) + if "" in dotted_path.split("."): + raise ConfigError(f"override {token} has an empty key in its path") + if not has_inline_value: + if position == len(argv): + raise ConfigError(f"override {token} is missing its value") + value = argv[position] + position += 1 + parsed: Any = value + if value == "null" or value[:1] in ("[", "{"): + try: + parsed = json.loads(value) + except (ValueError, RecursionError) as error: + raise ConfigError( + f"override {token} is not valid JSON: {error}" + ) from error + overrides.pop(dotted_path, None) # a repeated path applies where it is last + overrides[dotted_path] = parsed + return overrides + + +def apply_overrides( + document: Mapping[str, Any], overrides: Mapping[str, Any] +) -> dict[str, Any]: + """A copy of the parsed file with each override written at its dotted path, + in order. Mappings missing on the way are created and a parent that is not + a mapping (a scalar, a list, null) is replaced. A mapping value merges key + by key into a mapping already there, so `--model '{"depth": 3}'` keeps the + file's other `model` keys; a list or a scalar replaces (a repeated mapping + path from `parse_overrides` replaces too, since the mapping keeps one + value per path). Values are written as given for pydantic to validate. + Neither argument is written into: the copy takes dicts and lists apart, so + two keys sharing one object (a YAML anchor) stop sharing; a value that + contains itself is a `ConfigError`.""" + try: + copy = _copy_tree(document) + for dotted_path, value in overrides.items(): + _write(copy, dotted_path.split("."), value) + except RecursionError: + raise ConfigError( + "the config contains a value that refers to itself or is nested too deeply" + ) from None + return copy + + +def _write(mapping: dict[str, Any], keys: list[str], value: Any) -> None: + *parent_keys, last_key = keys + for key in parent_keys: # create mappings on the way; replace non-mappings + if not isinstance(mapping.get(key), dict): + mapping[key] = {} + mapping = mapping[key] + if isinstance(value, Mapping) and isinstance(mapping.get(last_key), dict): + for key, item in value.items(): # a JSON key is one key, dots included + _write(mapping[last_key], [key], item) + else: + mapping[last_key] = _copy_tree(value) + + +def _copy_tree(value: Any) -> Any: + """Copy the dicts and lists of a document, scalars as they are. Unlike + `copy.deepcopy`, two keys sharing one object get two copies.""" + if isinstance(value, Mapping): + return {key: _copy_tree(item) for key, item in value.items()} + if isinstance(value, list): + return [_copy_tree(item) for item in value] + return value From 7207b014ea567eaebb01268ca9279dc4ebecc188 Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Mon, 28 Sep 2026 12:18:23 +0200 Subject: [PATCH 10/14] Drop E0s and parents from the model metadata, freeze config sections, share the test schema (CORE-2 follow-up, #1556) Decisions from the reforge meeting: - ModelMetadata no longer records a head's E0s (E0Details is gone; the head's parameters in the model hold them, so nothing is lost) nor the records of the models it was built from (ParentModel and `parents` are gone for now). HeadSummary keeps the sources a head consumed. - Provenance records `versions`, a version per distribution involved (`{"mace-core": ..., "mace-torch": ...}`), in place of the single `code_version` of mace-core: the packages version independently, so one number does not identify the code. - A data source's `elements` are atomic numbers, not chemical symbols. - ReforgeBaseConfig is BaseConfig: the class name outlives the project name. - A ConfigSection is frozen once validated: what validation produced is what the run uses, and a change is a new validation of an edited dict (`from_dict`), never an assignment behind it. `frozen=True` cannot be reopened by a subclass, like `extra="forbid"`. - The two config test modules defined the same demo schema twice; it now lives once in tests/mace_core_demo.py (the union of both), and the three CLI tests that only repeated a base test through parse/apply are folded into the composed override test. Coverage of mace_core is unchanged (98%, the same five uncovered lines: the tomli fallback, the uninstalled-version fallback and two citation-rendering branches). - The exports' JSON mode and their inf setting are now pinned under a free field; either could be removed with every test passing. The config tests no longer promise that NaN survives: the metadata refuses to store one. - The config docstrings and the README entry are cut down to what the code does not say by itself. --- packages/mace-core/README.md | 23 +-- packages/mace-core/src/mace_core/__init__.py | 4 +- .../src/mace_core/config/__init__.py | 4 +- .../mace-core/src/mace_core/config/base.py | 63 ++++--- .../mace-core/src/mace_core/config/cli.py | 12 +- packages/mace-core/src/mace_core/metadata.py | 61 ++----- packages/mace-core/tests/mace_core_demo.py | 95 +++++++++++ .../mace-core/tests/test_mace_core_config.py | 154 ++++++------------ .../tests/test_mace_core_config_cli.py | 101 +----------- .../tests/test_mace_core_config_kinds.py | 8 +- .../tests/test_mace_core_metadata.py | 72 +++----- 11 files changed, 237 insertions(+), 360 deletions(-) create mode 100644 packages/mace-core/tests/mace_core_demo.py diff --git a/packages/mace-core/README.md b/packages/mace-core/README.md index 6039c2c0e..4f36ee2a7 100644 --- a/packages/mace-core/README.md +++ b/packages/mace-core/README.md @@ -9,20 +9,13 @@ Distribution `mace-core`, import name `mace_core`. What is here so far: -- `mace_core.config` — `ReforgeBaseConfig`, the pydantic base every v1 - config schema derives from: `load()` reads one TOML/YAML/JSON file and - validates it once (`from_dict()` for a caller that edits the parsed dict - first); unknown keys are pydantic errors at every level; - `to_resolved_dict()` is the fully defaulted, round-trippable export and - `to_user_dict()` holds only what was set. The base has no command-line - knowledge; `config.cli` holds the `--a.b value` grammar (`parse_overrides()` - to a mapping of dotted paths, `apply_overrides()` to write it into the - parsed dict) for the CLI to compose with `read_config_file()` and - `from_dict()`. +- `mace_core.config` — `BaseConfig`, the pydantic base every v1 config schema + derives from: one TOML/YAML/JSON file, validated once; unknown keys are + errors, a validated config is frozen, and its resolved export loads back to + the same config. `config.cli` holds the `--a.b value` override grammar for a + command line to compose with it. - `mace_core.metadata` — `ModelMetadata`, the versioned record every trained model carries (config as written and as resolved, provenance, a summary per - data source, per head its E0s and the sources it consumed, DOI, citations, - notes, and the records of the models it was built from), with a JSON round - trip, - `ConfigRecord.from_config()` to embed a config in its fixed-point form, and - `format_citations()`. + data source, per head the sources it consumed, DOI, citations and notes), + with a JSON round trip, `ConfigRecord.from_config()` to embed a config in + its fixed-point form, and `format_citations()`. diff --git a/packages/mace-core/src/mace_core/__init__.py b/packages/mace-core/src/mace_core/__init__.py index 882339eb6..c6920d9ae 100644 --- a/packages/mace-core/src/mace_core/__init__.py +++ b/packages/mace-core/src/mace_core/__init__.py @@ -6,14 +6,14 @@ from importlib.metadata import PackageNotFoundError, version -from mace_core.config import ConfigError, ConfigSection, ReforgeBaseConfig +from mace_core.config import BaseConfig, ConfigError, ConfigSection from mace_core.metadata import ModelMetadata, format_citations __all__ = [ + "BaseConfig", "ConfigError", "ConfigSection", "ModelMetadata", - "ReforgeBaseConfig", "__version__", "format_citations", ] diff --git a/packages/mace-core/src/mace_core/config/__init__.py b/packages/mace-core/src/mace_core/config/__init__.py index c9178ccfa..d80e5fd27 100644 --- a/packages/mace-core/src/mace_core/config/__init__.py +++ b/packages/mace-core/src/mace_core/config/__init__.py @@ -6,17 +6,17 @@ """ from mace_core.config.base import ( + BaseConfig, ConfigError, ConfigSection, - ReforgeBaseConfig, read_config_file, ) from mace_core.config.cli import apply_overrides, parse_overrides __all__ = [ + "BaseConfig", "ConfigError", "ConfigSection", - "ReforgeBaseConfig", "apply_overrides", "parse_overrides", "read_config_file", diff --git a/packages/mace-core/src/mace_core/config/base.py b/packages/mace-core/src/mace_core/config/base.py index 534662d32..b1aba12de 100644 --- a/packages/mace-core/src/mace_core/config/base.py +++ b/packages/mace-core/src/mace_core/config/base.py @@ -1,16 +1,8 @@ -"""One config file, one validation, two exports. - -`ReforgeBaseConfig.load(config_file)` reads one TOML, YAML or JSON file into a -dict (`read_config_file`, public) and validates it once with pydantic -(`from_dict`, public, for a caller that edits the dict first: `cli` does, for a -command line). Everything the schema rejects, unknown keys included, is -pydantic's `ValidationError` at its dotted location; `ConfigError` is raised -only for a file that cannot be read or parsed. `to_resolved_dict` is the full -JSON-native dump, itself a reloadable config file; `to_user_dict` holds only -what was set. A kinds field (a discriminated union on `kind`) is an ordinary -pydantic feature written in the tagged form, `loss: {kind: huber, delta: 0.1}`; -nothing here knows about it. Nothing here knows about a command line either; -the `--a.b value` grammar is `cli`'s. +"""The config base: one file, validated once by pydantic, exported as a dict. + +Everything the schema rejects, unknown keys included, is pydantic's +`ValidationError`; `ConfigError` is only for a file that cannot be read or +parsed. Nothing here knows about a command line: that is `cli`. """ import json @@ -30,7 +22,7 @@ else: import tomli as tomllib -__all__ = ["ConfigError", "ConfigSection", "ReforgeBaseConfig", "read_config_file"] +__all__ = ["BaseConfig", "ConfigError", "ConfigSection", "read_config_file"] class ConfigError(ValueError): @@ -82,21 +74,21 @@ def read_config_file(path: str | os.PathLike[str]) -> dict[str, Any]: class ConfigSection(BaseModel): - """A node of a config tree: unknown keys are errors at every level. - - Subclasses declare fields only. The definition-time checks keep unknown - keys fatal (every reachable model is a section, `extra="forbid"` cannot be - reopened) and the resolved export a fixed point (no sets, aliases, - excluded or computed fields). A class left incomplete by a forward - reference is checked on every `load` of a root that reaches it instead; - a forward reference to a class local to a function cannot be resolved - that way (pydantic cannot see the function's scope), so declare such - classes at module level. + """A node of a config tree: unknown keys are errors, and its attributes + cannot be reassigned (a list or dict it holds is not protected). A change + is a new validation, `from_dict` on an edited dict. + + Subclasses declare fields only. A subclass that would undo either rule, or + hold what does not load back from the export (a set, an alias, an excluded + or computed field, a plain `BaseModel`), is a `TypeError` when it is + defined, or on `load` if a forward reference left it incomplete. Declare a + forward-referenced class at module level: pydantic cannot resolve one + local to a function. """ # inf/nan as the JSON constants Infinity and NaN, not pydantic's default # null, which would turn a value into a different one. - model_config = ConfigDict(extra="forbid", ser_json_inf_nan="constants") + model_config = ConfigDict(extra="forbid", frozen=True, ser_json_inf_nan="constants") @classmethod def __pydantic_init_subclass__(cls, **kwargs: Any) -> None: @@ -107,6 +99,11 @@ def __pydantic_init_subclass__(cls, **kwargs: Any) -> None: f"{cls.__name__} sets extra={extra!r}; a section keeps " "extra='forbid' so an unknown key stays an error" ) + if not cls.model_config.get("frozen"): + raise TypeError( + f"{cls.__name__} sets frozen=False; a section stays frozen " + "so a validated config is not changed behind the validation" + ) if cls.__pydantic_complete__: # else a forward reference: `load` checks it _check_field_declarations(cls) @@ -167,8 +164,9 @@ def _check_sections_reached_by(root: type[ConfigSection]) -> None: to_visit.append(leaf) -class ReforgeBaseConfig(ConfigSection): - """The root of a config tree: `load` a file, or `from_dict` a parsed one.""" +class BaseConfig(ConfigSection): + """Base class for the root of a config schema: subclass it and declare the + fields, then `load` a file or `from_dict` a parsed one.""" @classmethod def load(cls, config_file: str | os.PathLike[str]) -> Self: @@ -177,11 +175,10 @@ def load(cls, config_file: str | os.PathLike[str]) -> Self: @classmethod def from_dict(cls, document: Mapping[str, Any]) -> Self: - """Build the config from a parsed document, checking the sections the - class reaches first. This is `load` for a caller that edits the parsed - file before validating it, such as a command line writing its flags into - the dict `read_config_file` returned. The document is not written into - (values under an `Any`-typed field are shared with it, not copied).""" + """Build the config from a parsed document: `load` for a caller that + edits the dict `read_config_file` returned first, such as a command + line. The document is not written into (values under an `Any`-typed + field are shared with it, not copied).""" _check_sections_reached_by(cls) return cls.model_validate(document) @@ -192,7 +189,7 @@ def to_resolved_dict(self) -> dict[str, Any]: return self.model_dump(mode="json") def to_user_dict(self) -> dict[str, Any]: - """Only what the file set, in the same shape. A loaded config carries + """Only what was set, in the same shape. A loaded config carries the tag of every kinds field it wrote, so this loads back; a variant built in code without its tag (`Config(loss=Huber(delta=2))`) exports without `kind`: pass the tag, or use `to_resolved_dict`.""" diff --git a/packages/mace-core/src/mace_core/config/cli.py b/packages/mace-core/src/mace_core/config/cli.py index 48c9becd8..d3ec3b5cb 100644 --- a/packages/mace-core/src/mace_core/config/cli.py +++ b/packages/mace-core/src/mace_core/config/cli.py @@ -66,14 +66,10 @@ def apply_overrides( ) -> dict[str, Any]: """A copy of the parsed file with each override written at its dotted path, in order. Mappings missing on the way are created and a parent that is not - a mapping (a scalar, a list, null) is replaced. A mapping value merges key - by key into a mapping already there, so `--model '{"depth": 3}'` keeps the - file's other `model` keys; a list or a scalar replaces (a repeated mapping - path from `parse_overrides` replaces too, since the mapping keeps one - value per path). Values are written as given for pydantic to validate. - Neither argument is written into: the copy takes dicts and lists apart, so - two keys sharing one object (a YAML anchor) stop sharing; a value that - contains itself is a `ConfigError`.""" + a mapping is replaced. A mapping value merges key by key into a mapping + already there, so `--model '{"depth": 3}'` keeps the file's other `model` + keys; anything else replaces. Neither argument is written into; a value + that contains itself is a `ConfigError`.""" try: copy = _copy_tree(document) for dotted_path, value in overrides.items(): diff --git a/packages/mace-core/src/mace_core/metadata.py b/packages/mace-core/src/mace_core/metadata.py index 575732c99..31066a975 100644 --- a/packages/mace-core/src/mace_core/metadata.py +++ b/packages/mace-core/src/mace_core/metadata.py @@ -18,7 +18,7 @@ from pydantic import BaseModel, ConfigDict, Field, model_validator -from mace_core.config import ReforgeBaseConfig +from mace_core.config import BaseConfig __all__ = [ "SCHEMA_VERSION", @@ -26,11 +26,9 @@ "ConfigRecord", "DataSourceSummary", "DataSummary", - "E0Details", "HeadSummary", "MetadataSchemaError", "ModelMetadata", - "ParentModel", "Provenance", "format_citations", ] @@ -55,7 +53,7 @@ class _Record(BaseModel): class ConfigRecord(_Record): """The training configuration, as written and as resolved. - Both are the JSON-native dicts a `ReforgeBaseConfig` exports: `user` is + Both are the JSON-native dicts a `BaseConfig` exports: `user` is `to_user_dict()`, the keys the config file and the command line set, and `resolved` is `to_resolved_dict()`, every key with defaults filled in. Build it with `from_config()` so the two cannot be mixed up. @@ -65,15 +63,17 @@ class ConfigRecord(_Record): resolved: dict[str, Any] = Field(default_factory=dict) @classmethod - def from_config(cls, config: ReforgeBaseConfig) -> ConfigRecord: + def from_config(cls, config: BaseConfig) -> ConfigRecord: return cls(user=config.to_user_dict(), resolved=config.to_resolved_dict()) class Provenance(_Record): """Which code produced the model.""" - #: Version of the `mace-core` distribution, as `mace_core.__version__` reports it. - code_version: str + #: Version per distribution involved, `{"mace-core": "1.0.2", "mace-torch": + #: "1.1.0"}`: the packages version independently, so no single number + #: identifies the code. + versions: dict[str, str] #: Full hash of the commit the code was run from; None when not in a checkout. git_commit: str | None = None @@ -91,8 +91,8 @@ class DataSourceSummary(_Record): name: str num_configurations: int | None = None num_atoms: int | None = None - #: Chemical symbols of every element present. - elements: list[str] = Field(default_factory=list) + #: Atomic numbers of every element present. + elements: list[int] = Field(default_factory=list) reference_keys: list[str] = Field(default_factory=list) @@ -103,27 +103,10 @@ class DataSummary(_Record): sources: list[DataSourceSummary] = Field(default_factory=list) -class E0Details(_Record): - """How one head's per-element reference energies (the E0s) were obtained. - - `values` maps chemical symbol to E0 in the model's energy unit. Symbols - rather than atomic numbers, because JSON keys are strings and an integer - key would not survive the round trip. - """ - - #: "explicit" when the E0s were given; "estimated" when fitted from the data. - source: Literal["explicit", "estimated"] - #: The estimation method (e.g. "average", "least_squares"); None when explicit. - method: str | None = None - #: Parameters of the method or reference, e.g. which data the fit used. - parameters: dict[str, Any] = Field(default_factory=dict) - values: dict[str, float] = Field(default_factory=dict) - - class HeadSummary(_Record): - """What one head was fitted on: its E0s and the data sources it consumed.""" + """What one head was fitted on: the data sources it consumed. Its E0s are + not recorded here; the head's parameters in the model hold them.""" - e0: E0Details #: Names in `DataSummary.sources`; a source feeding two heads appears in both. sources: list[str] = Field(default_factory=list) @@ -153,9 +136,6 @@ class ModelMetadata(_Record): doi: str | None = None citations: list[Citation] = Field(default_factory=list) notes: str = "" - #: The models this one was built from, each with its own record inside, so - #: the whole lineage travels with the model. - parents: list[ParentModel] = Field(default_factory=list) @model_validator(mode="after") def _heads_name_known_sources(self) -> ModelMetadata: @@ -191,9 +171,7 @@ def from_json(cls, text: str) -> ModelMetadata: """Parse a record written by `to_json()`. Raises `MetadataSchemaError` when the record carries a schema version - this code does not read, before any field is interpreted. Embedded - parent records are validated as fields, so a version mismatch inside - one is pydantic's error; a record only embeds parents it could read. + this code does not read, before any field is interpreted. """ try: document = json.loads(text) @@ -224,21 +202,6 @@ def from_json(cls, text: str) -> ModelMetadata: return cls.model_validate_json(text) -class ParentModel(_Record): - """A model this one was built from.""" - - #: What the parent contributed: its weights as the starting point (fine-tuning - #: or continued training), or its predictions as distillation targets. - role: Literal["initial_weights", "teacher"] - #: How the config named it: a path or a registry name such as "mace-mp-0b3". - name: str - #: The parent's own record; None for a legacy checkpoint that carries none. - metadata: ModelMetadata | None = None - - -ModelMetadata.model_rebuild() # `parents` refers to the class defined after it - - def format_citations(citations: Iterable[Citation]) -> str: """Render citations as a numbered, printable block; empty for none. diff --git a/packages/mace-core/tests/mace_core_demo.py b/packages/mace-core/tests/mace_core_demo.py new file mode 100644 index 000000000..faa6c294e --- /dev/null +++ b/packages/mace-core/tests/mace_core_demo.py @@ -0,0 +1,95 @@ +"""The demo schema the config tests share: two levels of nesting, a list, an +optional, a free dict, a kinds field; one config as a dict; and the helpers +that write it out in each format and read a `ValidationError`'s locations.""" + +import json +from typing import Annotated, Any, Literal + +import yaml +from mace_core.config import BaseConfig, ConfigSection +from pydantic import Field + + +class RadialSection(ConfigSection): + num_bessel: int = 8 + cutoff: float = 5.0 + + +class ModelSection(ConfigSection): + num_interactions: int = 2 + hidden_irreps: str = "128x0e + 128x1o" + radial: RadialSection = RadialSection() + + +class DataSection(ConfigSection): + train_file: str | None = None + valid_fraction: float = 0.1 + energy_key: str = "REF_energy" + heads: list[str] = Field(default_factory=lambda: ["default"]) + + +class StageTwoSection(ConfigSection): + start_epoch: int = 100 + energy_weight: float = 1000.0 + + +class Weighted(ConfigSection): + kind: Literal["weighted"] = "weighted" + stress_weight: float = 0.0 + + +class Huber(ConfigSection): + kind: Literal["huber"] = "huber" + delta: float = 0.01 + + +class DemoConfig(BaseConfig): + name: str = "mace" + seed: int = 123 + default_dtype: str = "float64" + model: ModelSection = ModelSection() + data: DataSection = DataSection() + #: A section left at its defaults unless a file or the CLI writes into it. + stage_two: StageTwoSection = StageTwoSection() + loss: Annotated[Weighted | Huber, Field(discriminator="kind")] = Weighted() + extra: dict[str, Any] = Field(default_factory=dict) + + +#: One config, as a dict. Each format test writes it out and loads it back. +FILE_VALUES = { + "name": "water", + "seed": 7, + "model": {"num_interactions": 4, "radial": {"cutoff": 4.5}}, + "data": {"train_file": "train.xyz", "heads": ["pbe", "r2scan"]}, +} + + +def to_toml(values, prefix=""): + """Enough TOML for a None-free config: scalars and lists share JSON's + literal syntax, nested dicts become `[a.b]` tables after the scalars.""" + lines = [ + f"{k} = {json.dumps(v)}" for k, v in values.items() if not isinstance(v, dict) + ] + for key, value in values.items(): + if isinstance(value, dict): + lines += [f"\n[{prefix}{key}]", to_toml(value, f"{prefix}{key}.")] + return "\n".join(lines) + + +def dump(values, extension): + if extension == ".toml": + return to_toml(values) + if extension == ".json": + return json.dumps(values) + return yaml.safe_dump(values) + + +def write_config(tmp_path, extension=".json", values=FILE_VALUES, name="config"): + path = tmp_path / f"{name}{extension}" + path.write_text(dump(values, extension), encoding="utf-8") + return path + + +def error_locations(excinfo): + """The dotted location of every error a `ValidationError` carries.""" + return [".".join(map(str, error["loc"])) for error in excinfo.value.errors()] diff --git a/packages/mace-core/tests/test_mace_core_config.py b/packages/mace-core/tests/test_mace_core_config.py index c9414d477..c5950b8c1 100644 --- a/packages/mace-core/tests/test_mace_core_config.py +++ b/packages/mace-core/tests/test_mace_core_config.py @@ -1,4 +1,4 @@ -"""`ReforgeBaseConfig`: file formats, `from_dict`, unknown keys, and the two +"""`BaseConfig`: file formats, `from_dict`, unknown keys, and the two exports' fixed point.""" import inspect @@ -12,93 +12,21 @@ import pytest import yaml -from mace_core.config import ( - ConfigError, - ConfigSection, - ReforgeBaseConfig, - read_config_file, +from mace_core.config import BaseConfig, ConfigError, ConfigSection, read_config_file +from mace_core_demo import ( + FILE_VALUES, + DemoConfig, + RadialSection, + StageTwoSection, + dump, + error_locations, + write_config, ) from pydantic import BaseModel, ConfigDict, Field, ValidationError, computed_field #: A warning the test did not ask for is a failure. pytestmark = pytest.mark.filterwarnings("error") -# --------------------------------------------------------------------------- -# The demo schema: two levels of nesting, a list, an optional, a Literal. - - -class RadialSection(ConfigSection): - num_bessel: int = 8 - cutoff: float = 5.0 - - -class ModelSection(ConfigSection): - num_interactions: int = 2 - hidden_irreps: str = "128x0e + 128x1o" - radial: RadialSection = RadialSection() - - -class DataSection(ConfigSection): - train_file: str | None = None - valid_fraction: float = 0.1 - energy_key: str = "REF_energy" - heads: list[str] = Field(default_factory=lambda: ["default"]) - - -class StageTwoSection(ConfigSection): - start_epoch: int = 100 - energy_weight: float = 1000.0 - - -class DemoConfig(ReforgeBaseConfig): - name: str = "mace" - seed: int = 123 - default_dtype: str = "float64" - model: ModelSection = ModelSection() - data: DataSection = DataSection() - #: A section left at its defaults unless a file or the CLI writes into it. - stage_two: StageTwoSection = StageTwoSection() - - -#: One config, as a dict. Each format test writes it out and loads it back. -FILE_VALUES = { - "name": "water", - "seed": 7, - "model": {"num_interactions": 4, "radial": {"cutoff": 4.5}}, - "data": {"train_file": "train.xyz", "heads": ["pbe", "r2scan"]}, -} - - -def to_toml(values, prefix=""): - """Enough TOML for a None-free config: scalars and lists share JSON's - literal syntax, nested dicts become `[a.b]` tables after the scalars.""" - lines = [ - f"{k} = {json.dumps(v)}" for k, v in values.items() if not isinstance(v, dict) - ] - for key, value in values.items(): - if isinstance(value, dict): - lines += [f"\n[{prefix}{key}]", to_toml(value, f"{prefix}{key}.")] - return "\n".join(lines) - - -def dump(values, extension): - if extension == ".toml": - return to_toml(values) - if extension == ".json": - return json.dumps(values) - return yaml.safe_dump(values) - - -def write_config(tmp_path, extension, values=FILE_VALUES, name="config"): - path = tmp_path / f"{name}{extension}" - path.write_text(dump(values, extension), encoding="utf-8") - return path - - -def error_locations(excinfo): - """The dotted location of every error a `ValidationError` carries.""" - return [".".join(map(str, error["loc"])) for error in excinfo.value.errors()] - # --------------------------------------------------------------------------- # File loading @@ -175,7 +103,7 @@ def test_malformed_file_is_a_config_error(tmp_path, extension, text): def test_a_yaml_anchor_that_contains_itself_is_a_config_error(tmp_path): # Under a section it would be pydantic's error; under a free dict it would # load and then fail to export, so the file is refused up front. - class Free(ReforgeBaseConfig): + class Free(BaseConfig): extra: dict[str, Any] = Field(default_factory=dict) path = tmp_path / "loop.yaml" @@ -279,7 +207,7 @@ def test_unknown_key_in_a_document_is_reported_at_its_path(): def test_every_bad_list_item_is_reported(tmp_path): - class Layers(ReforgeBaseConfig): + class Layers(BaseConfig): layers: list[RadialSection] = Field(default_factory=list) path = write_config(tmp_path, ".json", {"layers": [{"cutof": 1}, {"nb": 2}]}) @@ -309,6 +237,8 @@ def test_resolved_dict_has_every_default_in_declaration_order(tmp_path): "model", "data", "stage_two", + "loss", + "extra", ] assert resolved["stage_two"] == {"start_epoch": 100, "energy_weight": 1000.0} assert list(resolved["model"]) == ["num_interactions", "hidden_irreps", "radial"] @@ -342,37 +272,38 @@ def test_fixed_point_holds_through_toml_when_nothing_is_none(tmp_path): assert_fixed_point(tmp_path, first, extension) -def test_inf_and_nan_survive_the_exports_in_every_format(tmp_path): - # pydantic's JSON mode writes them as null by default, which would put a - # different value into the model metadata. Each format spells them its own - # way (JSON constants, YAML .inf/.nan, TOML inf/nan); nan != nan, so the - # round trip is compared as JSON text. +def test_exports_hold_json_native_values_under_a_free_field(): + # Without the JSON mode of the dump the tuple would stay a tuple. + config = DemoConfig.from_dict({"extra": {"window": (1, 2)}}) + assert config.to_resolved_dict()["extra"] == {"window": [1, 2]} + assert config.to_user_dict() == {"extra": {"window": [1, 2]}} + + +def test_inf_survives_the_exports_in_every_format(tmp_path): + # pydantic's JSON mode writes it as null by default under a free field + # (`extra`), which would put a different value into the model metadata. + # Each format spells it its own way: JSON Infinity, YAML .inf, TOML inf. values = { **FILE_VALUES, # the file sets the optional file name, so no null "model": {"radial": {"cutoff": math.inf}}, - "stage_two": {"energy_weight": math.nan}, + "extra": {"limit": math.inf}, } config = DemoConfig.load(write_config(tmp_path, ".yaml", values)) resolved = config.to_resolved_dict() assert resolved["model"]["radial"]["cutoff"] == math.inf - assert math.isnan(resolved["stage_two"]["energy_weight"]) - assert math.isnan(config.to_user_dict()["stage_two"]["energy_weight"]) + assert config.to_user_dict()["extra"] == {"limit": math.inf} text = json.dumps(resolved) - assert "null" not in text and "Infinity" in text and "NaN" in text + assert "null" not in text and "Infinity" in text for extension, body in [ (".json", text), (".yaml", yaml.safe_dump(resolved)), - (".toml", "[model.radial]\ncutoff = inf\n[stage_two]\nenergy_weight = nan\n"), + (".toml", "[model.radial]\ncutoff = inf\n"), ]: path = tmp_path / f"special{extension}" path.write_text(body, encoding="utf-8") second = DemoConfig.load(path).to_resolved_dict() assert second["model"]["radial"]["cutoff"] == math.inf, extension - assert math.isnan(second["stage_two"]["energy_weight"]), extension - assert ( - json.dumps(DemoConfig.load(tmp_path / "special.json").to_resolved_dict()) - == text - ) + assert DemoConfig.load(tmp_path / "special.json").to_resolved_dict() == resolved class LenientSection(BaseModel): @@ -412,13 +343,32 @@ def double(self) -> int: return 2 * self.seed -def test_a_section_cannot_reopen_extra(): +def test_a_section_cannot_reopen_extra_or_thaw(): with pytest.raises(TypeError, match=r"Loose sets extra='allow'; a section keeps"): class Loose(ConfigSection): model_config = ConfigDict(extra="allow") seed: int = 1 + with pytest.raises(TypeError, match=r"Thawed sets frozen=False; a section stays"): + + class Thawed(ConfigSection): + model_config = ConfigDict(frozen=False) + seed: int = 1 + + +def test_a_validated_config_is_immutable(tmp_path): + # What validation produced is what the run uses: a change is a new + # validation of an edited dict, never an assignment behind it. + config = DemoConfig.load(write_config(tmp_path, ".yaml")) + with pytest.raises(ValidationError, match="frozen"): + config.seed = 8 # ty: ignore[invalid-assignment] # the point of the test + with pytest.raises(ValidationError, match="frozen"): + config.model.radial.cutoff = 6.0 # ty: ignore[invalid-assignment] + assert config.seed == 7 and config.model.radial.cutoff == 4.5 + changed = DemoConfig.from_dict({**config.to_resolved_dict(), "seed": 8}) + assert changed.seed == 8 and changed.model == config.model + # A class that names a class defined below it is incomplete at definition: # pydantic keeps the name, so the check cannot see through the field. The @@ -433,7 +383,7 @@ class Later(ConfigSection): x: int = 1 -class ForwardConfig(ReforgeBaseConfig): +class ForwardConfig(BaseConfig): forward: Forward = Field(default_factory=Forward) @@ -445,7 +395,7 @@ class PlainLater(BaseModel): # not a ConfigSection: it would swallow a typo a: int = 1 -class LeakingConfig(ReforgeBaseConfig): +class LeakingConfig(BaseConfig): leaking: Leaking = Field(default_factory=Leaking) diff --git a/packages/mace-core/tests/test_mace_core_config_cli.py b/packages/mace-core/tests/test_mace_core_config_cli.py index ee7b1d112..0af28c504 100644 --- a/packages/mace-core/tests/test_mace_core_config_cli.py +++ b/packages/mace-core/tests/test_mace_core_config_cli.py @@ -4,93 +4,30 @@ `read_config_file` and `from_dict`.""" import json -from typing import Annotated, Any, Literal +from typing import Any import pytest from mace_core.config import ( ConfigError, - ConfigSection, - ReforgeBaseConfig, apply_overrides, parse_overrides, read_config_file, ) -from pydantic import Field, ValidationError +from mace_core_demo import FILE_VALUES, DemoConfig, Huber, error_locations, write_config +from pydantic import ValidationError #: A warning the test did not ask for is a failure. pytestmark = pytest.mark.filterwarnings("error") -# --------------------------------------------------------------------------- -# The demo schema: two levels of nesting, a list, an optional, a free dict, a -# kinds field. - - -class RadialSection(ConfigSection): - num_bessel: int = 8 - cutoff: float = 5.0 - - -class ModelSection(ConfigSection): - num_interactions: int = 2 - radial: RadialSection = RadialSection() - - -class DataSection(ConfigSection): - train_file: str | None = None - heads: list[str] = Field(default_factory=lambda: ["default"]) - - -class StageTwoSection(ConfigSection): - start_epoch: int = 100 - energy_weight: float = 1000.0 - - -class Weighted(ConfigSection): - kind: Literal["weighted"] = "weighted" - stress_weight: float = 0.0 - - -class Huber(ConfigSection): - kind: Literal["huber"] = "huber" - delta: float = 0.01 - - -class DemoConfig(ReforgeBaseConfig): - name: str = "mace" - seed: int = 123 - model: ModelSection = ModelSection() - data: DataSection = DataSection() - stage_two: StageTwoSection = StageTwoSection() - loss: Annotated[Weighted | Huber, Field(discriminator="kind")] = Weighted() - extra: dict[str, Any] = Field(default_factory=dict) - - -FILE_VALUES = { - "name": "water", - "seed": 7, - "model": {"num_interactions": 4, "radial": {"cutoff": 4.5}}, - "data": {"train_file": "train.xyz", "heads": ["pbe", "r2scan"]}, -} - HEADS_JSON = '["a", "b"]' -def write_config(tmp_path, values=FILE_VALUES): - path = tmp_path / "config.json" - path.write_text(json.dumps(values), encoding="utf-8") - return path - - def load(tmp_path, argv, values=FILE_VALUES, root=DemoConfig): """What a command line does with its config file and the tokens after it.""" - document = read_config_file(write_config(tmp_path, values)) + document = read_config_file(write_config(tmp_path, ".json", values)) return root.from_dict(apply_overrides(document, parse_overrides(argv))) -def error_locations(excinfo): - return [".".join(map(str, error["loc"])) for error in excinfo.value.errors()] - - # --------------------------------------------------------------------------- # parse_overrides: tokens to a mapping of dotted path to value. @@ -248,32 +185,12 @@ def test_an_override_beats_the_file_which_beats_the_defaults(tmp_path): assert config.model.radial.cutoff == 4.5 # the file's other values survive assert config.model.radial.num_bessel == 8 # defaults fill the rest assert load(tmp_path, []) == DemoConfig.from_dict(FILE_VALUES) - - -def test_values_are_handed_to_pydantic_as_given(tmp_path): - argv = ["--seed=9", "--data.train_file", "null", "--data.heads", HEADS_JSON] - config = load(tmp_path, argv, {}) - assert config.seed == 9 # pydantic's lax coercion - assert config.data.train_file is None - assert config.data.heads == ["a", "b"] - assert load(tmp_path, ["--stage_two.start_epoch", "50"], {}).stage_two == ( - StageTwoSection(start_epoch=50) - ) - - -def test_a_value_of_the_wrong_type_is_a_validation_error(tmp_path): - with pytest.raises(ValidationError, match="seed"): - load(tmp_path, ["--seed", "seven"]) - - -@pytest.mark.parametrize( - "dotted_path", ["nmae", "model.num_interaction", "stage_two.start"] -) -def test_an_unknown_key_in_an_override_is_reported_at_its_path(tmp_path, dotted_path): + # A section the file left alone is created on the way and validated like + # any other; the value stays the string the grammar handed over. + assert load(tmp_path, ["--stage_two.start_epoch", "50"]).stage_two.start_epoch == 50 with pytest.raises(ValidationError) as excinfo: - load(tmp_path, [f"--{dotted_path}", "1"], {}) - assert error_locations(excinfo) == [dotted_path] - assert excinfo.value.errors()[0]["type"] == "extra_forbidden" + load(tmp_path, ["--stage_two.start", "50"]) + assert error_locations(excinfo) == ["stage_two.start"] def test_an_override_writes_into_a_kinds_field(tmp_path): diff --git a/packages/mace-core/tests/test_mace_core_config_kinds.py b/packages/mace-core/tests/test_mace_core_config_kinds.py index 36c5c0c1a..f23029bf5 100644 --- a/packages/mace-core/tests/test_mace_core_config_kinds.py +++ b/packages/mace-core/tests/test_mace_core_config_kinds.py @@ -9,7 +9,7 @@ from typing import Annotated, Literal import pytest -from mace_core.config import ConfigSection, ReforgeBaseConfig +from mace_core.config import BaseConfig, ConfigSection from pydantic import Field, ValidationError #: A warning the test did not ask for is a failure. @@ -65,7 +65,7 @@ class HeadSection(ConfigSection): loss: Choice = Weighted() -class LossConfig(ReforgeBaseConfig): +class LossConfig(BaseConfig): energy_weight: float = 1.0 choice: Choice = Weighted() opt: OptChoice = NoChoice() @@ -79,7 +79,7 @@ class HuberRequired(ConfigSection): delta: float = 0.01 -class RequiredConfig(ReforgeBaseConfig): +class RequiredConfig(BaseConfig): choice: Annotated[Weighted | HuberRequired, Field(discriminator="kind")] = ( Weighted() ) @@ -250,7 +250,7 @@ def test_json_schema_is_produced_in_both_modes(): def test_a_variant_class_as_a_plain_field_keeps_kind_as_a_key(tmp_path): - class Reuse(ReforgeBaseConfig): + class Reuse(BaseConfig): direct: Huber = Huber() resolved = { diff --git a/packages/mace-core/tests/test_mace_core_metadata.py b/packages/mace-core/tests/test_mace_core_metadata.py index 4fa877f02..8ea12bc84 100644 --- a/packages/mace-core/tests/test_mace_core_metadata.py +++ b/packages/mace-core/tests/test_mace_core_metadata.py @@ -5,18 +5,16 @@ import sys import pytest -from mace_core.config import ConfigSection, ReforgeBaseConfig +from mace_core.config import BaseConfig, ConfigSection from mace_core.metadata import ( SCHEMA_VERSION, Citation, ConfigRecord, DataSourceSummary, DataSummary, - E0Details, HeadSummary, MetadataSchemaError, ModelMetadata, - ParentModel, Provenance, format_citations, ) @@ -39,33 +37,25 @@ def full_record() -> ModelMetadata: user={"model": {"num_interactions": 3}}, resolved={"name": "mace", "model": {"num_interactions": 3, "cutoff": 5.0}}, ), - provenance=Provenance(code_version="1.0.0", git_commit="a" * 40), + provenance=Provenance( + versions={"mace-core": "1.0.2", "mace-torch": "1.1.0"}, + git_commit="a" * 40, + ), data=DataSummary( sources=[ DataSourceSummary( name="water", num_configurations=1200, num_atoms=64_000, - elements=["H", "O"], + elements=[1, 8], reference_keys=["pbe_energy", "pbe_forces"], ), - DataSourceSummary(name="ice", elements=["H", "O"]), + DataSourceSummary(name="ice", elements=[1, 8]), ] ), heads={ - "pbe": HeadSummary( - e0=E0Details( - source="estimated", - method="least_squares", - parameters={"reference_key": "pbe_energy"}, - values={"H": -13.6, "O": -430.2}, - ), - sources=["water", "ice"], - ), - "r2scan": HeadSummary( - e0=E0Details(source="explicit", values={"H": -13.7, "O": -431.0}), - sources=["ice"], - ), + "pbe": HeadSummary(sources=["water", "ice"]), + "r2scan": HeadSummary(sources=["ice"]), }, doi="10.5281/zenodo.0000000", citations=[MACE_PAPER, Citation(title="A dataset paper", doi="10.1000/xyz")], @@ -84,7 +74,7 @@ def test_json_round_trip_is_lossless(): def test_minimal_record_round_trips_too(): record = ModelMetadata( - config=ConfigRecord(), provenance=Provenance(code_version="0.0.0") + config=ConfigRecord(), provenance=Provenance(versions={"mace-core": "0.0.0"}) ) assert ModelMetadata.from_json(record.to_json()) == record assert record.heads == {} @@ -103,35 +93,16 @@ def test_heads_must_name_summarised_sources(): def test_config_and_provenance_are_mandatory(): with pytest.raises(ValidationError, match="config"): - ModelMetadata.model_validate({"provenance": {"code_version": "0"}}) + ModelMetadata.model_validate({"provenance": {"versions": {"mace-core": "0"}}}) def test_lossy_value_is_refused_rather_than_stored(): record = full_record() - record.heads["pbe"].e0.parameters["shape"] = (2, 3) # JSON brings it back as a list + record.config.user["shape"] = (2, 3) # JSON brings it back as a list with pytest.raises(MetadataSchemaError, match="does not survive a JSON round trip"): record.to_json() -def test_lineage_round_trips_through_two_levels(): - foundation = ParentModel(role="initial_weights", name="mace-mp-0b3") # no record - distilled = full_record() - distilled.parents = [ - foundation, - ParentModel(role="teacher", name="teacher.model", metadata=full_record()), - ] - fine_tuned = full_record() - fine_tuned.parents = [ - ParentModel(role="initial_weights", name="distilled.model", metadata=distilled) - ] - back = ModelMetadata.from_json(fine_tuned.to_json()) - assert back == fine_tuned - assert back.parents[0].metadata is not None - grandparents = back.parents[0].metadata.parents - assert [p.role for p in grandparents] == ["initial_weights", "teacher"] - assert grandparents[0].metadata is None - - def test_schema_version_is_written(): assert json.loads(full_record().to_json())["schema_version"] == SCHEMA_VERSION @@ -175,14 +146,14 @@ def test_non_record_json_is_rejected_with_context(text, message): def test_infinity_survives_and_nan_is_refused(): - # pydantic's default writes inf/nan as null, which would silently turn an - # E0 into a different value. NaN is never equal to itself, so it cannot - # pass the round-trip check; an E0 or a config value that is NaN is a bug + # pydantic's default writes inf/nan as null, which would silently turn a + # config value into a different one. NaN is never equal to itself, so it + # cannot pass the round-trip check; a config value that is NaN is a bug # upstream, not something to store. record = full_record() - record.heads["pbe"].e0.values["H"] = float("inf") + record.config.resolved["model"]["cutoff"] = float("inf") back = ModelMetadata.from_json(record.to_json()) - assert back.heads["pbe"].e0.values["H"] == float("inf") + assert back.config.resolved["model"]["cutoff"] == float("inf") record.config.resolved["cutoff"] = float("nan") with pytest.raises(MetadataSchemaError, match="does not survive"): record.to_json() @@ -198,20 +169,15 @@ def test_schema_version_is_pinned_on_direct_validation_as_well(): def test_unknown_fields_are_rejected(): with pytest.raises(ValidationError, match="extra_forbidden"): ModelMetadata.model_validate( - {"config": {}, "provenance": {"code_version": "0"}, "note": "x"} + {"config": {}, "provenance": {"versions": {"mace-core": "0"}}, "note": "x"} ) -def test_e0_source_is_one_of_two_values(): - with pytest.raises(ValidationError, match="source"): - E0Details.model_validate({"source": "guessed"}) - - def test_config_record_is_built_from_a_config(): class Section(ConfigSection): cutoff: float = 5.0 - class Config(ReforgeBaseConfig): + class Config(BaseConfig): seed: int = 1 model: Section = Section() From 8ed74a463a9cddf1619011d76bfad05de21c4b86 Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Thu, 1 Oct 2026 13:10:58 +0200 Subject: [PATCH 11/14] Review base.py: drop the YAML cycle check, name the schema checks for what they do (CORE-2 follow-up, #1556) - read_config_file no longer refuses a YAML anchor that contains itself. Nobody writes one, and without the check it still fails loudly: a pydantic ValidationError under a typed section, a ValueError at export under a free field. - _leaf_types is _types_in: it returned every class an annotation mentions, containers included, and now filters out non-classes itself. - _check_field_declarations is _check_section_fields and _check_sections_reached_by is _check_schema; their docstrings say why the check runs at definition and again on load (a forward reference leaves a section unchecked until then). - The docstrings of from_dict, to_resolved_dict and to_user_dict are one sentence each. --- .../mace-core/src/mace_core/config/base.py | 84 +++++++++---------- .../mace-core/tests/test_mace_core_config.py | 16 +--- 2 files changed, 39 insertions(+), 61 deletions(-) diff --git a/packages/mace-core/src/mace_core/config/base.py b/packages/mace-core/src/mace_core/config/base.py index b1aba12de..3d065f66f 100644 --- a/packages/mace-core/src/mace_core/config/base.py +++ b/packages/mace-core/src/mace_core/config/base.py @@ -33,9 +33,8 @@ class ConfigError(ValueError): def read_config_file(path: str | os.PathLike[str]) -> dict[str, Any]: """Parse one TOML, YAML or JSON config file, chosen by its (case-folded) extension, into a dict; an empty or comment-only file is `{}`. Raises - `ConfigError` for an unknown extension, an unreadable file, a parse failure, - a top level that is not a mapping, or a value that contains itself (a YAML - anchor inside itself), which no schema could export again.""" + `ConfigError` for an unknown extension, an unreadable file, a parse failure + or a top level that is not a mapping.""" path = Path(path) parsers = { ".toml": tomllib.loads, @@ -64,12 +63,6 @@ def read_config_file(path: str | os.PathLike[str]) -> dict[str, Any]: f"config file {path} must be a mapping of keys to values at the top " f"level, not {type(document).__name__}" ) - try: # the stdlib's cycle detector; shared siblings pass, only a cycle fails - json.dumps(document, default=str, skipkeys=True) - except ValueError as error: - raise ConfigError( - f"config file {path} contains a value that refers to itself" - ) from error return document @@ -104,21 +97,25 @@ def __pydantic_init_subclass__(cls, **kwargs: Any) -> None: f"{cls.__name__} sets frozen=False; a section stays frozen " "so a validated config is not changed behind the validation" ) - if cls.__pydantic_complete__: # else a forward reference: `load` checks it - _check_field_declarations(cls) + # A forward reference leaves the field types unknown for now; + # `_check_schema` checks the section on load instead. + if cls.__pydantic_complete__: + _check_section_fields(cls) -def _leaf_types(annotation: Any) -> Iterator[Any]: - """Every class or origin an annotation reaches through `Annotated`, unions and - list, tuple or dict parameters. Non-type arguments (Literal values, `Field` - metadata) come out too; callers test `isinstance(leaf, type)`.""" - yield get_origin(annotation) or annotation +def _types_in(annotation: Any) -> Iterator[type]: + """Every class an annotation mentions, containers included.""" + outer = get_origin(annotation) or annotation + if isinstance(outer, type): + yield outer for argument in get_args(annotation): - yield from _leaf_types(argument) + yield from _types_in(argument) -def _check_field_declarations(section: type[ConfigSection]) -> None: - """The checks on one class's own fields; a violation is a `TypeError`.""" +def _check_section_fields(section: type[ConfigSection]) -> None: + """Refuse, with a `TypeError`, what pydantic allows on a field but a config + section cannot have: a field that would not load back from the export, or + a child that would ignore unknown keys.""" if section.model_computed_fields: name = next(iter(section.model_computed_fields)) raise TypeError( @@ -130,25 +127,25 @@ def _check_field_declarations(section: type[ConfigSection]) -> None: raise TypeError(f"{where} has an alias; a config key is its field name") if field.exclude: raise TypeError(f"{where} is excluded from dumps; it would not load back") - for leaf in _leaf_types(field.annotation): - if not isinstance(leaf, type): - continue + for held in _types_in(field.annotation): # Any set type, `set`, `frozenset` or an abstract one, dumps in an # order that varies between runs, so the export would not be stable. - if issubclass(leaf, AbstractSet): + if issubclass(held, AbstractSet): raise TypeError(f"{where} is typed as a set; order varies. Use a list") - if issubclass(leaf, BaseModel) and not issubclass(leaf, ConfigSection): + if issubclass(held, BaseModel) and not issubclass(held, ConfigSection): raise TypeError( - f"{where} holds {leaf.__name__}, which is not a ConfigSection; " + f"{where} holds {held.__name__}, which is not a ConfigSection; " "unknown keys under it would be dropped" ) -def _check_sections_reached_by(root: type[ConfigSection]) -> None: - """Resolve any section the root's tree reaches that a forward reference left - incomplete at definition, and check every reached section. Checking again on - each load is a few attribute reads per class; it saves remembering which - classes were checked.""" +def _check_schema(root: type[ConfigSection]) -> None: + """Run `_check_section_fields` on the root and every section under it. + + A section is checked when it is defined, unless a forward reference left + its field types unknown then. Such a section is resolved and checked here, + on load. The others are checked a second time, which costs less than + tracking which ones were skipped.""" to_visit, seen = [root], set() while to_visit: section = to_visit.pop() @@ -157,11 +154,11 @@ def _check_sections_reached_by(root: type[ConfigSection]) -> None: seen.add(section) if not section.__pydantic_complete__: section.model_rebuild() # resolves the forward reference or raises - _check_field_declarations(section) + _check_section_fields(section) for field in section.model_fields.values(): - for leaf in _leaf_types(field.annotation): - if isinstance(leaf, type) and issubclass(leaf, ConfigSection): - to_visit.append(leaf) + for held in _types_in(field.annotation): + if issubclass(held, ConfigSection): + to_visit.append(held) class BaseConfig(ConfigSection): @@ -175,22 +172,17 @@ def load(cls, config_file: str | os.PathLike[str]) -> Self: @classmethod def from_dict(cls, document: Mapping[str, Any]) -> Self: - """Build the config from a parsed document: `load` for a caller that - edits the dict `read_config_file` returned first, such as a command - line. The document is not written into (values under an `Any`-typed - field are shared with it, not copied).""" - _check_sections_reached_by(cls) + """Build the config from an already parsed file, such as one a command + line applied its overrides to.""" + _check_schema(cls) return cls.model_validate(document) def to_resolved_dict(self) -> dict[str, Any]: - """Every field, defaults filled, as JSON-native values in declaration - order: a config file that loads back to this config (through TOML - whenever no value is `None`).""" + """A JSON-compatible dict that loads back to this config, defaults + filled in.""" return self.model_dump(mode="json") def to_user_dict(self) -> dict[str, Any]: - """Only what was set, in the same shape. A loaded config carries - the tag of every kinds field it wrote, so this loads back; - a variant built in code without its tag (`Config(loss=Huber(delta=2))`) - exports without `kind`: pass the tag, or use `to_resolved_dict`.""" + """A JSON-compatible dict that loads back to this config, holding only + the fields that were set.""" return self.model_dump(mode="json", exclude_unset=True) diff --git a/packages/mace-core/tests/test_mace_core_config.py b/packages/mace-core/tests/test_mace_core_config.py index c5950b8c1..420b44562 100644 --- a/packages/mace-core/tests/test_mace_core_config.py +++ b/packages/mace-core/tests/test_mace_core_config.py @@ -8,7 +8,7 @@ import subprocess import sys from collections.abc import Set as AbstractSet -from typing import Annotated, Any +from typing import Annotated import pytest import yaml @@ -100,20 +100,6 @@ def test_malformed_file_is_a_config_error(tmp_path, extension, text): DemoConfig.load(path) -def test_a_yaml_anchor_that_contains_itself_is_a_config_error(tmp_path): - # Under a section it would be pydantic's error; under a free dict it would - # load and then fail to export, so the file is refused up front. - class Free(BaseConfig): - extra: dict[str, Any] = Field(default_factory=dict) - - path = tmp_path / "loop.yaml" - path.write_text("extra: &loop {b: *loop}\n", encoding="utf-8") - with pytest.raises(ConfigError, match=r"loop\.yaml contains a value that refers"): - Free.load(path) - path.write_text("extra: {a: &shared {x: 1}, b: *shared}\n", encoding="utf-8") - assert Free.load(path).extra == {"a": {"x": 1}, "b": {"x": 1}} # sharing is fine - - # --------------------------------------------------------------------------- # Precedence: defaults < the file. Nothing else feeds a config. From a338e99532a5c9e2ffd3706a1436efbc3ecd691e Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Thu, 1 Oct 2026 14:02:00 +0200 Subject: [PATCH 12/14] Review cli.py: drop the checks the schema already makes, split the path write from the merge (CORE-2 follow-up, #1556) - parse_overrides no longer refuses an empty key in a path (`--a..b`): it reaches pydantic as the unknown key "" and fails there. Only under a free dict field is it stored, where it shows in the export. - A value that contains itself, or JSON nested past the recursion limit, is no longer turned into a ConfigError, in apply_overrides, parse_overrides and read_config_file alike: nobody writes one, and it still fails, as a RecursionError. - _write did two things and is two functions: _set_at_path walks the dotted path and _set_or_merge sets the value or merges a dict into a dict already there. - The docstrings of parse_overrides and apply_overrides open with one sentence that says what the function does. --- .../mace-core/src/mace_core/config/base.py | 2 +- .../mace-core/src/mace_core/config/cli.py | 58 +++++++++---------- .../tests/test_mace_core_config_cli.py | 16 ----- 3 files changed, 30 insertions(+), 46 deletions(-) diff --git a/packages/mace-core/src/mace_core/config/base.py b/packages/mace-core/src/mace_core/config/base.py index 3d065f66f..3fef486ac 100644 --- a/packages/mace-core/src/mace_core/config/base.py +++ b/packages/mace-core/src/mace_core/config/base.py @@ -54,7 +54,7 @@ def read_config_file(path: str | os.PathLike[str]) -> dict[str, Any]: raise ConfigError(f"cannot read config file {path}: {error}") from error try: document = parse(text) - except (ValueError, yaml.YAMLError, RecursionError) as error: # toml, json: Value + except (ValueError, yaml.YAMLError) as error: # toml, json: ValueError raise ConfigError(f"cannot parse config file {path}: {error}") from error if document is None: # YAML reads an empty or comment-only file as None return {} diff --git a/packages/mace-core/src/mace_core/config/cli.py b/packages/mace-core/src/mace_core/config/cli.py index d3ec3b5cb..da3f08aeb 100644 --- a/packages/mace-core/src/mace_core/config/cli.py +++ b/packages/mace-core/src/mace_core/config/cli.py @@ -22,12 +22,12 @@ def parse_overrides(tokens: Iterable[str]) -> dict[str, Any]: - """`--a.b.c value` or `--a.b.c=value`, in order, to `{"a.b.c": value}`. A - value that is `null` or starts with `[` or `{` is parsed as JSON; any other - value stays a string for pydantic to coerce. A repeated path keeps its last - value, at its last position. A token that is not `--path`, a path without - its value or with an empty key (`--a..b`), and a JSON value that does not - parse are `ConfigError`s.""" + """Turn command-line tokens into a dict of dotted path to value. + + Both `--a.b value` and `--a.b=value` give `{"a.b": value}`. A value that + is `null` or starts with `[` or `{` is parsed as JSON; any other stays a + string for pydantic to coerce. A repeated path keeps its last value. A + token that fits none of this is a `ConfigError`.""" if isinstance(tokens, str): raise TypeError(f"tokens is a string, {tokens!r}; pass a list of tokens") argv = list(tokens) @@ -41,8 +41,6 @@ def parse_overrides(tokens: Iterable[str]) -> dict[str, Any]: raise ConfigError( f"unknown config option '{token}'; options are --key.path value" ) - if "" in dotted_path.split("."): - raise ConfigError(f"override {token} has an empty key in its path") if not has_inline_value: if position == len(argv): raise ConfigError(f"override {token} is missing its value") @@ -52,7 +50,7 @@ def parse_overrides(tokens: Iterable[str]) -> dict[str, Any]: if value == "null" or value[:1] in ("[", "{"): try: parsed = json.loads(value) - except (ValueError, RecursionError) as error: + except ValueError as error: raise ConfigError( f"override {token} is not valid JSON: {error}" ) from error @@ -64,34 +62,36 @@ def parse_overrides(tokens: Iterable[str]) -> dict[str, Any]: def apply_overrides( document: Mapping[str, Any], overrides: Mapping[str, Any] ) -> dict[str, Any]: - """A copy of the parsed file with each override written at its dotted path, - in order. Mappings missing on the way are created and a parent that is not - a mapping is replaced. A mapping value merges key by key into a mapping - already there, so `--model '{"depth": 3}'` keeps the file's other `model` - keys; anything else replaces. Neither argument is written into; a value - that contains itself is a `ConfigError`.""" - try: - copy = _copy_tree(document) - for dotted_path, value in overrides.items(): - _write(copy, dotted_path.split("."), value) - except RecursionError: - raise ConfigError( - "the config contains a value that refers to itself or is nested too deeply" - ) from None + """Return a copy of a parsed config file with the overrides written in. + + Each value goes to its dotted path, in order. A mapping merges into a + mapping already there, so `--model '{"depth": 3}'` keeps the file's other + `model` keys; any other value replaces what was there.""" + copy = _copy_tree(document) + for dotted_path, value in overrides.items(): + _set_at_path(copy, dotted_path.split("."), value) return copy -def _write(mapping: dict[str, Any], keys: list[str], value: Any) -> None: +def _set_at_path(mapping: dict[str, Any], keys: list[str], value: Any) -> None: + """Walk down the keys, creating dicts on the way (a parent that is not a + dict is replaced by one), and set or merge the value at the last key.""" *parent_keys, last_key = keys - for key in parent_keys: # create mappings on the way; replace non-mappings + for key in parent_keys: if not isinstance(mapping.get(key), dict): mapping[key] = {} mapping = mapping[key] - if isinstance(value, Mapping) and isinstance(mapping.get(last_key), dict): - for key, item in value.items(): # a JSON key is one key, dots included - _write(mapping[last_key], [key], item) + _set_or_merge(mapping, last_key, value) + + +def _set_or_merge(mapping: dict[str, Any], key: str, value: Any) -> None: + """Set `mapping[key]`. A dict merges, key by key, into a dict already + there; anything else replaces what was there.""" + if isinstance(value, Mapping) and isinstance(mapping.get(key), dict): + for inner_key, item in value.items(): + _set_or_merge(mapping[key], inner_key, item) else: - mapping[last_key] = _copy_tree(value) + mapping[key] = _copy_tree(value) def _copy_tree(value: Any) -> Any: diff --git a/packages/mace-core/tests/test_mace_core_config_cli.py b/packages/mace-core/tests/test_mace_core_config_cli.py index 0af28c504..18a07f5d4 100644 --- a/packages/mace-core/tests/test_mace_core_config_cli.py +++ b/packages/mace-core/tests/test_mace_core_config_cli.py @@ -4,7 +4,6 @@ `read_config_file` and `from_dict`.""" import json -from typing import Any import pytest from mace_core.config import ( @@ -57,12 +56,6 @@ def test_an_empty_inline_value_does_not_hide_the_next_option(): assert parse_overrides(["--name=", "--seed", "5"]) == {"name": "", "seed": "5"} -@pytest.mark.parametrize("token", ["--a..b", "--a.", "--.a"]) -def test_an_empty_key_in_a_path_is_a_config_error(token): - with pytest.raises(ConfigError, match=rf"override {token} has an empty key"): - parse_overrides([token, "1"]) - - def test_a_value_starting_with_dashes_works_in_both_forms(): assert parse_overrides(["--name=--odd"]) == {"name": "--odd"} assert parse_overrides(["--name", "--odd"]) == {"name": "--odd"} @@ -165,15 +158,6 @@ def test_a_yaml_anchor_does_not_share_an_override(tmp_path): assert apply_overrides(document, {"a.x": 1}) == {"a": {"x": 1}, "b": {}} -def test_a_value_that_contains_itself_is_a_config_error(): - loop: dict[str, Any] = {} - loop["b"] = loop - with pytest.raises(ConfigError, match="refers to itself"): - apply_overrides({"a": loop}, {}) - with pytest.raises(ConfigError, match="refers to itself"): - apply_overrides({}, {"a": loop}) - - # --------------------------------------------------------------------------- # Composed with the base: the override beats the file, which beats the # defaults, and every schema error is pydantic's at the path the override named. From 60eb06add1457e0b7bae5752cbb1dee0b5ff27ac Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Thu, 1 Oct 2026 14:32:52 +0200 Subject: [PATCH 13/14] Review metadata.py: one place for the schema version, from_json checks only the version (CORE-2 follow-up, #1556) - The schema version was written twice, in SCHEMA_VERSION and in the `Literal[1]` of the field. The field is now an int defaulting to the constant, so a bump is one edit. Only from_json checks the version; to_json still refuses a wrong one because it reads its own output back. Validating a dict directly no longer checks it. - from_json keeps the one check nothing else makes, the version, ahead of validation so that a newer record reports its version and not its unknown fields. Invalid JSON is json's own error, and a missing or non-integer version fails the same comparison. A top level that is not an object is now an AttributeError; to_json never writes one. - The round-trip error of to_json names what does not survive (a tuple, NaN) instead of the docstring. - The docstrings are cut down; the reference-key convention sits above `reference_keys`. --- packages/mace-core/src/mace_core/metadata.py | 113 +++++------------- .../tests/test_mace_core_metadata.py | 33 +---- 2 files changed, 34 insertions(+), 112 deletions(-) diff --git a/packages/mace-core/src/mace_core/metadata.py b/packages/mace-core/src/mace_core/metadata.py index 31066a975..d91b9690a 100644 --- a/packages/mace-core/src/mace_core/metadata.py +++ b/packages/mace-core/src/mace_core/metadata.py @@ -1,20 +1,11 @@ -"""The record every v1 model carries about how it was made. - -`ModelMetadata` is stored alongside the weights of every trained model, not -only foundation models. It is plain data: a Pydantic tree that serialises to -JSON with `to_json()` and comes back, without loss, through `from_json()`. - -The record is versioned. `SCHEMA_VERSION` is bumped whenever a field is -added, removed or changes meaning, and `from_json()` refuses a record written -under a version this code does not know, so a newer checkpoint fails loudly -at load time instead of being read with the wrong meanings. -""" +"""`ModelMetadata`: the record every trained v1 model carries about how it was +made, stored as JSON beside the weights.""" from __future__ import annotations import json from collections.abc import Iterable -from typing import Any, Final, Literal +from typing import Any, Final from pydantic import BaseModel, ConfigDict, Field, model_validator @@ -33,13 +24,13 @@ "format_citations", ] -#: The schema version this module writes and the only one it reads. Bump it -#: together with the `Literal` on `ModelMetadata.schema_version`. +#: Bump when a field is added, removed or changes meaning. SCHEMA_VERSION: Final = 1 class MetadataSchemaError(ValueError): - """The metadata was written under a schema version this code cannot read.""" + """Metadata this code cannot read back: an unknown schema version, or a + value JSON would change.""" class _Record(BaseModel): @@ -51,13 +42,7 @@ class _Record(BaseModel): class ConfigRecord(_Record): - """The training configuration, as written and as resolved. - - Both are the JSON-native dicts a `BaseConfig` exports: `user` is - `to_user_dict()`, the keys the config file and the command line set, and - `resolved` is `to_resolved_dict()`, every key with defaults filled in. - Build it with `from_config()` so the two cannot be mixed up. - """ + """Record of the training config, as written and as resolved.""" user: dict[str, Any] = Field(default_factory=dict) resolved: dict[str, Any] = Field(default_factory=dict) @@ -70,22 +55,14 @@ def from_config(cls, config: BaseConfig) -> ConfigRecord: class Provenance(_Record): """Which code produced the model.""" - #: Version per distribution involved, `{"mace-core": "1.0.2", "mace-torch": - #: "1.1.0"}`: the packages version independently, so no single number - #: identifies the code. + #: Version per distribution used, `{"mace-core": "1.0.2", "mace-torch": "1.1.0"}` versions: dict[str, str] #: Full hash of the commit the code was run from; None when not in a checkout. git_commit: str | None = None class DataSourceSummary(_Record): - """Automated summary of one data source. - - Reference-quantity keys name the method that produced the reference as a - prefix on the quantity: `pbe_energy`, `pbe_forces`, `r2scan_energy`. - `reference_keys` lists the keys this source provides, under that - convention. The heads a source fed name it in `ModelMetadata.heads`. - """ + """Summary of one data source.""" #: The data source's name in the config. name: str @@ -93,21 +70,21 @@ class DataSourceSummary(_Record): num_atoms: int | None = None #: Atomic numbers of every element present. elements: list[int] = Field(default_factory=list) + #: The reference quantities provided, each prefixed by the method that + #: produced it: `pbe_energy`, `pbe_forces`, `r2scan_energy`. reference_keys: list[str] = Field(default_factory=list) class DataSummary(_Record): - """One summary per data source, each once even when several heads share - it (`ModelMetadata` checks that); totals are sums over them, not stored.""" + """One summary per data source, each source once.""" sources: list[DataSourceSummary] = Field(default_factory=list) class HeadSummary(_Record): - """What one head was fitted on: the data sources it consumed. Its E0s are - not recorded here; the head's parameters in the model hold them.""" + """What one head was fitted on.""" - #: Names in `DataSummary.sources`; a source feeding two heads appears in both. + #: Names from `DataSummary.sources`. sources: list[str] = Field(default_factory=list) @@ -123,10 +100,9 @@ class Citation(_Record): class ModelMetadata(_Record): - """The mandatory per-model record. See the module docstring.""" + """How a model was made: its config, the code, the data and what to cite.""" - #: Pinned to the version this code reads; a bump here is a schema change. - schema_version: Literal[1] = SCHEMA_VERSION + schema_version: int = SCHEMA_VERSION config: ConfigRecord provenance: Provenance data: DataSummary = Field(default_factory=DataSummary) @@ -141,10 +117,7 @@ class ModelMetadata(_Record): def _heads_name_known_sources(self) -> ModelMetadata: names = [source.name for source in self.data.sources] if len(set(names)) != len(names): - raise ValueError( - f"data.sources names a source twice: {sorted(names)}; " - f"summarise each source once" - ) + raise ValueError(f"data.sources names a source twice: {sorted(names)}") for head, summary in self.heads.items(): for name in summary.sources: if name not in names: @@ -155,59 +128,35 @@ def _heads_name_known_sources(self) -> ModelMetadata: return self def to_json(self, indent: int | None = 2) -> str: - """Serialise; raises `MetadataSchemaError` if the text reads back to a - different record, so a lossy value (a tuple that comes back as a list, - a datetime that comes back as a string) is never stored.""" + """Serialise to JSON, checking that it reads back as the same record.""" text = self.model_dump_json(indent=indent) if self.from_json(text) != self: raise MetadataSchemaError( - "model metadata does not survive a JSON round trip; " - "a field holds a value JSON cannot represent" + "model metadata does not survive a JSON round trip: a field holds " + "a value JSON reads back differently, such as a tuple (a list " + "on the way back) or NaN (never equal to itself)" ) return text @classmethod def from_json(cls, text: str) -> ModelMetadata: - """Parse a record written by `to_json()`. - - Raises `MetadataSchemaError` when the record carries a schema version - this code does not read, before any field is interpreted. - """ - try: - document = json.loads(text) - except ValueError as error: - raise MetadataSchemaError( - f"model metadata is not valid JSON: {error}" - ) from error - if not isinstance(document, dict): - raise MetadataSchemaError( - f"model metadata must be a JSON object, not {type(document).__name__}" - ) + """Read a record written by `to_json`.""" + document = json.loads(text) + # Checked before validation: a newer record may hold fields this code + # does not know, and the version is the error worth reporting. version = document.get("schema_version") - if type(version) is not int: - raise MetadataSchemaError( - f"model metadata has schema_version {version!r}; expected the " - f"integer {SCHEMA_VERSION}" - ) if version != SCHEMA_VERSION: - hint = ( - "it was written by a newer mace-core; upgrade to read it" - if version > SCHEMA_VERSION - else "no migration exists for it" - ) raise MetadataSchemaError( - f"model metadata has schema_version {version}, but this " - f"mace-core reads schema_version {SCHEMA_VERSION}: {hint}" + f"model metadata has schema_version {version!r}, but this " + f"mace-core reads schema_version {SCHEMA_VERSION}; a newer " + "record needs an upgrade of mace-core" ) - return cls.model_validate_json(text) + return cls.model_validate(document) def format_citations(citations: Iterable[Citation]) -> str: - """Render citations as a numbered, printable block; empty for none. - - One line per citation: authors, title, venue and year, then the DOI or - URL. Fields that are unset are left out rather than printed as None. - """ + """Render citations as a numbered block, one line each; unset fields are + left out.""" lines = [] for number, citation in enumerate(citations, start=1): parts = [] diff --git a/packages/mace-core/tests/test_mace_core_metadata.py b/packages/mace-core/tests/test_mace_core_metadata.py index 8ea12bc84..785f26f9e 100644 --- a/packages/mace-core/tests/test_mace_core_metadata.py +++ b/packages/mace-core/tests/test_mace_core_metadata.py @@ -118,33 +118,13 @@ def test_future_schema_version_is_rejected_clearly(): assert "upgrade" in message -@pytest.mark.parametrize( - ("mutate", "message"), - [ - ( - lambda d: d.pop("schema_version"), - "schema_version None; expected the integer 1", - ), - (lambda d: d.update(schema_version="1"), "schema_version '1'; expected"), - (lambda d: d.update(schema_version=1.0), "schema_version 1.0; expected"), - ], -) -def test_missing_or_non_integer_schema_version_is_rejected(mutate, message): +def test_a_record_without_its_schema_version_is_rejected(): document = json.loads(full_record().to_json()) - mutate(document) - with pytest.raises(MetadataSchemaError, match=message): + del document["schema_version"] + with pytest.raises(MetadataSchemaError, match="schema_version None"): ModelMetadata.from_json(json.dumps(document)) -@pytest.mark.parametrize( - ("text", "message"), - [("[1]", "must be a JSON object, not list"), ("{", "is not valid JSON")], -) -def test_non_record_json_is_rejected_with_context(text, message): - with pytest.raises(MetadataSchemaError, match=message): - ModelMetadata.from_json(text) - - def test_infinity_survives_and_nan_is_refused(): # pydantic's default writes inf/nan as null, which would silently turn a # config value into a different one. NaN is never equal to itself, so it @@ -159,13 +139,6 @@ def test_infinity_survives_and_nan_is_refused(): record.to_json() -def test_schema_version_is_pinned_on_direct_validation_as_well(): - document = json.loads(full_record().to_json()) - document["schema_version"] = SCHEMA_VERSION + 1 - with pytest.raises(ValidationError, match="schema_version"): - ModelMetadata.model_validate(document) - - def test_unknown_fields_are_rejected(): with pytest.raises(ValidationError, match="extra_forbidden"): ModelMetadata.model_validate( From 803492c56930e11a7e54429cef3a3ad28479ee95 Mon Sep 17 00:00:00 2001 From: arnon-1 Date: Thu, 1 Oct 2026 15:42:06 +0200 Subject: [PATCH 14/14] Check the metadata schema version on every route, not only in from_json (CORE-2 follow-up, #1556) The check moves from from_json to a before-validator on ModelMetadata, so the constructor, model_validate and model_validate_json all refuse a record of another version, and do so before any field is read: a newer record reports its version, not its unknown fields. - from_json is model_validate_json. Invalid JSON and a top level that is not an object are pydantic's errors again. - A wrong version is a pydantic ValidationError carrying the message; MetadataSchemaError is left for a value that does not survive JSON. - A record with no schema_version key reads as the current version: the validator cannot tell it from a record built in code. to_json always writes the key. --- packages/mace-core/src/mace_core/metadata.py | 31 ++++++++++--------- .../tests/test_mace_core_metadata.py | 20 ++++++------ 2 files changed, 27 insertions(+), 24 deletions(-) diff --git a/packages/mace-core/src/mace_core/metadata.py b/packages/mace-core/src/mace_core/metadata.py index d91b9690a..fbd96f3b1 100644 --- a/packages/mace-core/src/mace_core/metadata.py +++ b/packages/mace-core/src/mace_core/metadata.py @@ -3,7 +3,6 @@ from __future__ import annotations -import json from collections.abc import Iterable from typing import Any, Final @@ -29,8 +28,7 @@ class MetadataSchemaError(ValueError): - """Metadata this code cannot read back: an unknown schema version, or a - value JSON would change.""" + """Metadata that would not read back from its JSON as the same record.""" class _Record(BaseModel): @@ -113,6 +111,21 @@ class ModelMetadata(_Record): citations: list[Citation] = Field(default_factory=list) notes: str = "" + @model_validator(mode="before") + @classmethod + def _reads_only_its_own_schema_version(cls, data: Any) -> Any: + # Before the fields are read: a newer record may hold fields this code + # does not know, and the version is the error worth reporting. + if isinstance(data, dict): + version = data.get("schema_version", SCHEMA_VERSION) + if version != SCHEMA_VERSION: + raise ValueError( + f"model metadata has schema_version {version!r}, but this " + f"mace-core reads schema_version {SCHEMA_VERSION}; a newer " + "record needs an upgrade of mace-core" + ) + return data + @model_validator(mode="after") def _heads_name_known_sources(self) -> ModelMetadata: names = [source.name for source in self.data.sources] @@ -141,17 +154,7 @@ def to_json(self, indent: int | None = 2) -> str: @classmethod def from_json(cls, text: str) -> ModelMetadata: """Read a record written by `to_json`.""" - document = json.loads(text) - # Checked before validation: a newer record may hold fields this code - # does not know, and the version is the error worth reporting. - version = document.get("schema_version") - if version != SCHEMA_VERSION: - raise MetadataSchemaError( - f"model metadata has schema_version {version!r}, but this " - f"mace-core reads schema_version {SCHEMA_VERSION}; a newer " - "record needs an upgrade of mace-core" - ) - return cls.model_validate(document) + return cls.model_validate_json(text) def format_citations(citations: Iterable[Citation]) -> str: diff --git a/packages/mace-core/tests/test_mace_core_metadata.py b/packages/mace-core/tests/test_mace_core_metadata.py index 785f26f9e..a971f9bd8 100644 --- a/packages/mace-core/tests/test_mace_core_metadata.py +++ b/packages/mace-core/tests/test_mace_core_metadata.py @@ -107,24 +107,24 @@ def test_schema_version_is_written(): assert json.loads(full_record().to_json())["schema_version"] == SCHEMA_VERSION -def test_future_schema_version_is_rejected_clearly(): +@pytest.mark.parametrize( + "read", + [ModelMetadata.model_validate, lambda d: ModelMetadata.from_json(json.dumps(d))], + ids=["model_validate", "from_json"], +) +def test_another_schema_version_is_rejected_on_every_route(read): document = json.loads(full_record().to_json()) document["schema_version"] = SCHEMA_VERSION + 1 - with pytest.raises(MetadataSchemaError) as excinfo: - ModelMetadata.from_json(json.dumps(document)) + document["a_field_of_the_next_version"] = 1 + with pytest.raises(ValidationError) as excinfo: + read(document) + assert excinfo.value.error_count() == 1 # the version, not the unknown field message = str(excinfo.value) assert f"schema_version {SCHEMA_VERSION + 1}" in message assert f"reads schema_version {SCHEMA_VERSION}" in message assert "upgrade" in message -def test_a_record_without_its_schema_version_is_rejected(): - document = json.loads(full_record().to_json()) - del document["schema_version"] - with pytest.raises(MetadataSchemaError, match="schema_version None"): - ModelMetadata.from_json(json.dumps(document)) - - def test_infinity_survives_and_nan_is_refused(): # pydantic's default writes inf/nan as null, which would silently turn a # config value into a different one. NaN is never equal to itself, so it