From e48c17f988e7f382601725e24b2cf3c5b95b62a2 Mon Sep 17 00:00:00 2001 From: Claude Date: Fri, 18 Sep 2026 23:02:54 +0000 Subject: [PATCH 1/2] refactor: no signature in the package says Any, and a symbol table section that is not a mapping is refused The public doors take Mapping[str, object] rather than dict[str, Any], a raw value is an object until pydantic has read it, the two grammars hand back the node type their child walk names, and the exclusivity proof compares a cell with a literal of its own kind. Narrowing the symbol table's sections turned an AttributeError on a list or a string under dimensions: or names: into a SchemaError naming the rewrite. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01EEeM2YoAk4Xr2uWB5qwMsH --- src/math_spec/_expression_parser.py | 23 +++-- src/math_spec/_where_parser.py | 40 +++++--- src/math_spec/_yaml.py | 10 +- src/math_spec/advice.py | 4 +- src/math_spec/exclusivity.py | 134 ++++++++++++++++---------- src/math_spec/lowering.py | 5 +- src/math_spec/model.py | 55 +++++++---- src/math_spec/piecewise.py | 32 +++--- src/math_spec/typesetting/__init__.py | 30 ++++-- src/math_spec/typesetting/symbols.py | 28 ++++-- src/math_spec/validation.py | 7 +- tests/typesetting/test_symbols.py | 2 + 12 files changed, 237 insertions(+), 133 deletions(-) diff --git a/src/math_spec/_expression_parser.py b/src/math_spec/_expression_parser.py index 6d0eb602..db0d9e18 100644 --- a/src/math_spec/_expression_parser.py +++ b/src/math_spec/_expression_parser.py @@ -12,7 +12,7 @@ from dataclasses import dataclass, field from functools import lru_cache -from typing import TYPE_CHECKING, Any, Literal, assert_never, cast, get_args +from typing import TYPE_CHECKING, Literal, assert_never, cast, get_args import pyparsing as pp @@ -409,16 +409,17 @@ def _make_func_call(tokens: pp.ParseResults) -> FunctionCallNode: return FunctionCallNode(name=name, args=tuple(args), kwargs=kwargs) -def _make_left_assoc(tokens: pp.ParseResults) -> Any: +def _make_left_assoc(tokens: pp.ParseResults) -> ArithmeticNode: + result: ArithmeticNode result, *rest = tokens for op, right in zip(rest[::2], rest[1::2], strict=True): result = BinaryOperatorNode(op, result, right) return result -def _make_power(tokens: pp.ParseResults) -> Any: +def _make_power(tokens: pp.ParseResults) -> ArithmeticNode: """A base and at most one exponent — right-associative, since the exponent is itself a ``unary``.""" - items = list(tokens) + items: list[ArithmeticNode] = list(tokens) return items[0] if len(items) == 1 else BinaryOperatorNode('**', items[0], items[2]) @@ -458,20 +459,22 @@ def _too_deep(what: str, text: str, found: int | None, rewrite: str) -> str: return f'The {what} {measured}, past the {MAX_DEPTH} levels the language admits: {shown!r}\n{rewrite}' -def parse_text( +def parse_text[T]( grammar: pp.ParserElement, text: str, what: str, rewrite: Callable[[str, int], str | None], - child_of: Callable[[Any], tuple[Any, ...]], + child_of: Callable[[T], tuple[T, ...]], deep_rewrite: str, -) -> Any: +) -> T: """Parse the whole of *text* with *grammar*, or raise :class:`SchemaError` naming *what* failed to parse. *rewrite* is asked for the predictable mistake at the failure position; its sentence, if any, precedes the grammar's own complaint. A tree nesting past :data:`MAX_DEPTH`, measured through *child_of*, is refused with - *deep_rewrite* — and so is one the parser itself ran out of stack on. + *deep_rewrite* — and so is one the parser itself ran out of stack on. The + node comes back as the type *child_of* walks, which is the grammar's word + for what it builds. """ try: result = grammar.parse_string(text, parse_all=True) @@ -481,7 +484,7 @@ def parse_text( raise SchemaError(msg) from e except RecursionError: raise SchemaError(_too_deep(what, text, None, deep_rewrite)) from None - node = result[0] + node = cast('T', result[0]) found = depth(node, child_of) if found > MAX_DEPTH: raise SchemaError(_too_deep(what, text, found, deep_rewrite)) @@ -535,4 +538,4 @@ def parse_expression(text: str) -> ParsedNode: lone ``=``, ``^`` for power — is named with its rewrite before the grammar's own complaint. """ - return cast('ParsedNode', parse_text(_GRAMMAR, text, 'expression', _named_rewrite, children, _DEEP_REWRITE)) + return parse_text(_GRAMMAR, text, 'expression', _named_rewrite, children, _DEEP_REWRITE) diff --git a/src/math_spec/_where_parser.py b/src/math_spec/_where_parser.py index 809436e8..4e6b805d 100644 --- a/src/math_spec/_where_parser.py +++ b/src/math_spec/_where_parser.py @@ -16,17 +16,24 @@ from dataclasses import dataclass from functools import lru_cache -from typing import TYPE_CHECKING, Any, cast, get_args +from typing import TYPE_CHECKING, cast, get_args import pyparsing as pp -from math_spec._expression_parser import ARITHMETIC, NAME, children, parse_text -from math_spec.program import AndNode, BooleanLiteralNode, NotNode, OrNode, PredicateOperator, where_children +from math_spec._expression_parser import ARITHMETIC, NAME, ArithmeticNode, children, parse_text +from math_spec.program import ( + AndNode, + BooleanLiteralNode, + ConnectiveWhereNode, + NotNode, + OrNode, + PredicateOperator, + where_children, +) if TYPE_CHECKING: from collections.abc import Callable - from math_spec._expression_parser import ArithmeticNode from math_spec.program import WhereNode # --------------------------------------------------------------------------- @@ -84,6 +91,11 @@ class UnresolvedComparisonNode: #: leaves are still names the schema has not been asked about. UnresolvedWhereNode = UnresolvedNameNode | UnresolvedComparisonNode +#: Every node a parsed where string is built of: the connectives and literals, +#: the unresolved leaves, and the arithmetic and the two side nodes under a +#: comparison. What the depth measurement walks. +type _ParsedWhere = WhereNode | UnresolvedWhereNode | ArithmeticNode | ColumnNode | QuotedNode + # --------------------------------------------------------------------------- # Grammar @@ -136,14 +148,14 @@ def _build_where_grammar() -> pp.ParserElement: return where_expr -def _folder(node_type: type[AndNode] | type[OrNode]) -> Callable[[pp.ParseResults], Any]: +def _folder(node_type: type[AndNode] | type[OrNode]) -> Callable[[pp.ParseResults], WhereNode | UnresolvedWhereNode]: """A parse action left-folding a flat operator chain into *node_type*.""" - def fold(tokens: pp.ParseResults) -> Any: - items = list(tokens) - result: WhereNode | UnresolvedWhereNode = items[0] + def fold(tokens: pp.ParseResults) -> WhereNode | UnresolvedWhereNode: + items: list[WhereNode | UnresolvedWhereNode] = list(tokens) + result = items[0] for item in items[1:]: - result = node_type(cast('WhereNode', result), item) + result = node_type(cast('WhereNode', result), cast('WhereNode', item)) return result return fold @@ -179,15 +191,15 @@ def _named_rewrite(text: str, loc: int) -> str | None: ) -def _nested(node: Any) -> tuple[Any, ...]: +def _nested(node: _ParsedWhere) -> tuple[_ParsedWhere, ...]: """What a where string nests through: a connective's operands, and the arithmetic on a comparison's sides.""" if isinstance(node, UnresolvedComparisonNode): return (node.left, node.right) - if isinstance(node, UnresolvedNameNode | ColumnNode | QuotedNode): - return () - if isinstance(node, AndNode | OrNode | NotNode | BooleanLiteralNode): + if isinstance(node, ArithmeticNode): + return children(node) + if isinstance(node, ConnectiveWhereNode): return where_children(node) - return children(node) + return () @lru_cache(maxsize=4096) diff --git a/src/math_spec/_yaml.py b/src/math_spec/_yaml.py index 696175b6..a703900b 100644 --- a/src/math_spec/_yaml.py +++ b/src/math_spec/_yaml.py @@ -22,7 +22,7 @@ import re from pathlib import Path -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING import yaml @@ -64,7 +64,7 @@ def _check_duplicate_keys(node: yaml.Node, origin: str) -> None: two merge keys are two merges, which PyYAML accumulates. """ if isinstance(node, yaml.MappingNode): - seen: dict[Any, int] = {} + seen: dict[str, int] = {} pairs: list[tuple[yaml.Node, yaml.Node]] = node.value for key_node, value_node in pairs: line = key_node.start_mark.line + 1 @@ -89,12 +89,12 @@ def _check_duplicate_keys(node: yaml.Node, origin: str) -> None: _check_duplicate_keys(item, origin) -def read_yaml(path: Path | str) -> dict[str, Any]: +def read_yaml(path: Path | str) -> dict[str, object]: """Read *path* off disk and parse it, in YAML 1.2's reading of scalars.""" return parse_yaml(Path(path).read_text(encoding='utf-8'), str(path)) -def read_model(model: str | Path) -> dict[str, Any]: +def read_model(model: str | Path) -> dict[str, object]: """A model from a file or from its text — a newline decides which a ``str`` is. A :class:`~pathlib.Path` names a file, and so does a ``str`` with no @@ -117,7 +117,7 @@ def read_model(model: str | Path) -> dict[str, Any]: return parse_yaml(model, 'YAML text') -def parse_yaml(text: str, origin: str = '') -> dict[str, Any]: +def parse_yaml(text: str, origin: str = '') -> dict[str, object]: """Parse YAML *text* as a mapping of sections. Args: diff --git a/src/math_spec/advice.py b/src/math_spec/advice.py index 5524ae16..f8e7a5b8 100644 --- a/src/math_spec/advice.py +++ b/src/math_spec/advice.py @@ -17,14 +17,14 @@ from math_spec.program import At, GroupSum, walk if TYPE_CHECKING: + from collections.abc import Mapping from pathlib import Path - from typing import Any from math_spec.model import Spec from math_spec.program import Program -def advice(model: str | Path | dict[str, Any] | Spec | Program) -> tuple[Advice, ...]: +def advice(model: str | Path | Mapping[str, object] | Spec | Program) -> tuple[Advice, ...]: """Everything the language advises about *model* — never an error, decidable without data. Args: diff --git a/src/math_spec/exclusivity.py b/src/math_spec/exclusivity.py index a669eb93..8e676743 100644 --- a/src/math_spec/exclusivity.py +++ b/src/math_spec/exclusivity.py @@ -19,7 +19,7 @@ import math from dataclasses import dataclass from enum import Enum -from typing import TYPE_CHECKING, Any, Literal, assert_never, cast +from typing import TYPE_CHECKING, Literal, Protocol, Self, assert_never from math_spec.program import ( AndNode, @@ -128,6 +128,19 @@ class Special(Enum): #: What one subject's value is, in one cell. Cell = float | str | bool | int | datetime.date | Special +#: What a where comparison is written against: a number, a label, or a date. +#: ``position()`` counts in integers, which are numbers here. +type _Literal = float | str | datetime.date + + +class _Ordered(Protocol): + """A value that orders against its own kind — what one atom's truth in one cell asks of both sides.""" + + def __lt__(self, other: Self, /) -> bool: ... + def __le__(self, other: Self, /) -> bool: ... + def __gt__(self, other: Self, /) -> bool: ... + def __ge__(self, other: Self, /) -> bool: ... + @dataclass(frozen=True) class Subject: @@ -166,7 +179,7 @@ class _Grid: @classmethod def of(cls, masks: Iterable[Mask], dtypes: Mapping[str, DeclaredDtype]) -> _Grid: - values: dict[Subject, set[Any]] = {} + values: dict[Subject, set[_Literal]] = {} subjects: dict[int, Subject] = {} for mask in masks: for node in mask.atoms: @@ -187,7 +200,9 @@ def witness(self, cell: dict[Subject, Cell]) -> str: return ', '.join(f'{subject} is {_shown(subject, value)}' for subject, value in cell.items()) -def _observe(node: TypedPredicateNode, subject: Subject, values: set[Any], dtypes: Mapping[str, DeclaredDtype]) -> None: +def _observe( + node: TypedPredicateNode, subject: Subject, values: set[_Literal], dtypes: Mapping[str, DeclaredDtype] +) -> None: """Record what *node* says about its subject: a position, or a literal. ``position()`` converts the dimension to an integer, so an ordering over a @@ -242,14 +257,14 @@ def _subject_of(node: TypedPredicateNode) -> Subject: assert_never(node) -def _cells_for(subject: Subject, values: set[Any], dtypes: Mapping[str, DeclaredDtype]) -> list[Cell]: +def _cells_for(subject: Subject, values: set[_Literal], dtypes: Mapping[str, DeclaredDtype]) -> list[Cell]: """Every region *subject*'s value can sit in — ordinary values first. The order is the order :func:`_witness` searches, so a refusal names an absent value or an infinity only where nothing plainer is a witness. """ if subject.kind == 'rank': - return _rank_cells(subject, cast('set[int]', values)) + return _rank_cells(subject, {int(v) for v in _numbers(values)}) if subject.kind in ('relation_pair', 'variable'): return [True, False] dtype = dtypes.get(subject.name) @@ -262,10 +277,13 @@ def _cells_for(subject: Subject, values: set[Any], dtypes: Mapping[str, Declared raise Undecidable(msg) return [True, False, Special.NULL] numeric = _numeric(dtype, values) - dated = _dated(values) - cells: list[Cell] = list( - _ordered_cells(values, discrete=dated or dtype == 'int') if numeric or dated else _label_cells(values) - ) + cells: list[Cell] + if numeric: + cells = list(_numeric_cells(_numbers(values), discrete=dtype == 'int')) + elif _dated(values): + cells = list(_dated_cells({v for v in values if isinstance(v, datetime.date)})) + else: + cells = _label_cells(values) cells.extend(_absence_cells(subject, numeric=numeric)) return cells @@ -282,57 +300,64 @@ def _absence_cells(subject: Subject, *, numeric: bool) -> list[Cell]: return [Special.NULL, Special.NEG_INF, Special.POS_INF] if numeric else [Special.NULL] -def _numeric(dtype: DeclaredDtype | None, literals: set[Any]) -> bool: +def _numeric(dtype: DeclaredDtype | None, literals: set[_Literal]) -> bool: """Is this subject a magnitude? The declaration says so where it is known.""" if dtype is not None: return dtype in ('float', 'int') return bool(literals) and all(isinstance(value, int | float) and not isinstance(value, bool) for value in literals) -def _dated(literals: set[Any]) -> bool: +def _dated(literals: set[_Literal]) -> bool: return bool(literals) and all(isinstance(value, datetime.date) for value in literals) -def _ordered_cells(literals: set[Any], *, discrete: bool) -> list[Cell]: - """Each literal, and one representative of the gap on either side of it.""" +def _numbers(literals: set[_Literal]) -> set[float]: + """The literals of a magnitude, as the numbers the dtype rules guarantee they are.""" + numbers = {float(value) for value in literals if isinstance(value, int | float)} + assert len(numbers) == len(literals), 'a label or a date reached a magnitude; the dtype rules keep them apart' + return numbers + + +def _numeric_cells(literals: set[float], *, discrete: bool) -> list[float]: + """Each number, and one representative of the gap on either side of it. + + An ``int`` subject has a value between two literals only where the gap is + wider than one. + """ if not literals: return [0.0] values = sorted(literals) - step = _step(values[0]) - cells: list[Cell] = [values[0] - step] + cells = [values[0] - 1.0] for value, following in itertools.zip_longest(values, values[1:]): cells.append(value) if following is None: continue - between = _between(value, following, step, discrete=discrete) - if between is not None: - cells.append(between) - cells.append(values[-1] + step) + if not discrete: + cells.append((value + following) / 2.0) + elif following - value > 1.0: + cells.append(value + 1.0) + cells.append(values[-1] + 1.0) return cells -def _step(value: Any) -> Any: - """How far outside the named literals a representative has to sit.""" - if isinstance(value, datetime.datetime): - return datetime.timedelta(seconds=1) - if isinstance(value, datetime.date): - return datetime.timedelta(days=1) - return 1.0 +def _dated_cells(literals: set[datetime.date]) -> list[datetime.date]: + """Each date, and one representative of the gap on either side of it. - -def _between(value: Any, following: Any, step: Any, *, discrete: bool) -> Any | None: - """A value strictly between two literals, or ``None`` where the type admits none. - - A discrete subject — an ``int`` or a date — has one only where the gap is wider than one unit. + A date steps by a day and a datetime by a second, and there is a value + between two literals only where the gap is wider than one step. """ - if discrete: - return value + step if following - value > step else None - # pyrefly: ignore[no-any-return-implicit] -- declaring `Any` would silence this and stop - # saying that the discrete branch has nothing to return. - return (value + following) / 2.0 + values = sorted(literals) + step = datetime.timedelta(seconds=1) if isinstance(values[0], datetime.datetime) else datetime.timedelta(days=1) + cells = [values[0] - step] + for value, following in itertools.zip_longest(values, values[1:]): + cells.append(value) + if following is not None and following - value > step: + cells.append(value + step) + cells.append(values[-1] + step) + return cells -def _label_cells(literals: set[Any]) -> list[Cell]: +def _label_cells(literals: set[_Literal]) -> list[Cell]: """Every named label, and one standing for all the labels not named.""" return [*sorted(literals, key=str), Special.OTHER] @@ -434,8 +459,13 @@ def _atom(node: TypedPredicateNode, cell: dict[Subject, Cell], grid: _Grid) -> b assert_never(node) -def _compare(value: Cell, op: PredicateOperator, literal: Any) -> bool: - """One atom's truth in one cell. Both sides are already this cell's frame.""" +def _compare(value: Cell, op: PredicateOperator, literal: _Literal) -> bool: + """One atom's truth in one cell. Both sides are already this cell's frame. + + Raises: + AssertionError: The cell and the literal are of different kinds, which + the dtype rules keep apart before a mask is proved. + """ if isinstance(value, Special): if value is Special.OTHER: # a label none of the masks names sorts nowhere @@ -443,24 +473,28 @@ def _compare(value: Cell, op: PredicateOperator, literal: Any) -> bool: return op == '!=' msg = f'a label neither case names is ordered with {op!r} — compare labels with == or != instead' raise Undecidable(msg) - magnitude = math.inf if value is Special.POS_INF else -math.inf - return _ordered(magnitude, op, float(literal)) - if isinstance(value, int | float) and isinstance(literal, int | float) and not isinstance(value, bool): + value = math.inf if value is Special.POS_INF else -math.inf + if isinstance(value, int | float) and isinstance(literal, int | float): return _ordered(float(value), op, float(literal)) - return _ordered(value, op, literal) + if isinstance(value, str) and isinstance(literal, str): + return _ordered(value, op, literal) + if isinstance(value, datetime.date) and isinstance(literal, datetime.date): + return _ordered(value, op, literal) + msg = f'{value!r} is compared with {literal!r}, and the two are of different kinds' + raise AssertionError(msg) -def _ordered(left: Any, op: PredicateOperator, right: Any) -> bool: +def _ordered[T: _Ordered](left: T, op: PredicateOperator, right: T) -> bool: match op: case '==': - return bool(left == right) + return left == right case '!=': - return bool(left != right) + return left != right case '<': - return bool(left < right) + return left < right case '<=': - return bool(left <= right) + return left <= right case '>': - return bool(left > right) + return left > right case '>=': - return bool(left >= right) + return left >= right diff --git a/src/math_spec/lowering.py b/src/math_spec/lowering.py index 1399a754..afcfc599 100644 --- a/src/math_spec/lowering.py +++ b/src/math_spec/lowering.py @@ -39,9 +39,8 @@ from math_spec.validation import to_spec if TYPE_CHECKING: - from collections.abc import Callable + from collections.abc import Callable, Mapping from pathlib import Path - from typing import Any from math_spec.model import Spec, _ExpandedSpec @@ -59,7 +58,7 @@ def _none_of(masks: list[program.Mask]) -> program.Mask: return remainder -def to_program(spec: str | Path | dict[str, Any] | Spec | program.Program) -> program.Program: +def to_program(spec: str | Path | Mapping[str, object] | Spec | program.Program) -> program.Program: """*spec* as a :class:`~math_spec.program.Program` — the public door. Takes whatever you have: a YAML path, the YAML itself, a mapping, a loaded diff --git a/src/math_spec/model.py b/src/math_spec/model.py index 71e6b21d..0bff3fc3 100644 --- a/src/math_spec/model.py +++ b/src/math_spec/model.py @@ -13,7 +13,7 @@ import re from collections import Counter from functools import cached_property -from typing import TYPE_CHECKING, Annotated, Any, ClassVar, Literal, Self, cast, get_args, override +from typing import TYPE_CHECKING, Annotated, ClassVar, Literal, Self, cast, get_args, override from pydantic import ( BaseModel, @@ -36,7 +36,8 @@ if TYPE_CHECKING: from collections.abc import Iterable, Iterator - from pydantic import GetJsonSchemaHandler + from pydantic import GetJsonSchemaHandler, SerializerFunctionWrapHandler + from pydantic.config import ExtraValues from pydantic.json_schema import JsonSchemaValue from pydantic_core import CoreSchema @@ -58,7 +59,7 @@ class _StrictBlock(BaseModel): @model_validator(mode='before') @classmethod - def _reject_unknown_keys(cls, data: Any) -> Any: + def _reject_unknown_keys(cls, data: object) -> object: """Name the near-miss, which is what a typo actually needs. pydantic's own ``extra='forbid'`` is the backstop; this runs first @@ -262,7 +263,7 @@ class BoundsBlock(_StrictBlock): @field_validator('lower', 'upper', mode='before') @classmethod - def _a_number_or_a_name(cls, v: Any, info: ValidationInfo) -> Any: + def _a_number_or_a_name(cls, v: object, info: ValidationInfo) -> object: if isinstance(v, bool): msg = f'bounds.{info.field_name} is a boolean, and a bound is a number or a parameter name.' raise ValueError(msg) @@ -353,7 +354,7 @@ def _check_formals(self) -> MacroBlock: return self -def _number_is_an_expression(value: Any) -> Any: +def _number_is_an_expression(value: object) -> object: """``expression: 0`` is how a file writes a constant — YAML reads it as an int. Booleans are left to fail: ``true`` is not arithmetic, and an error naming @@ -415,7 +416,7 @@ class ExpressionBlock(_StrictBlock): @model_validator(mode='before') @classmethod - def _from_string(cls, data: Any) -> Any: + def _from_string(cls, data: object) -> object: return {'expression': data} if isinstance(data, str) else data @model_validator(mode='after') @@ -462,9 +463,9 @@ def __get_pydantic_json_schema__(cls, core_schema: CoreSchema, handler: GetJsonS return _also_written_as(core_schema, handler, {'type': 'string'}) @model_serializer - def _as_written(self) -> str | dict[str, Any]: + def _as_written(self) -> str | dict[str, object]: if self.cases: - written: dict[str, Any] = {'dims': list(self.dims or [])} + written: dict[str, object] = {'dims': list(self.dims or [])} if self.description is not None: written['description'] = self.description written['cases'] = {name: case.model_dump() for name, case in self.cases.items()} @@ -492,7 +493,7 @@ class PiecewiseLink(_StrictBlock): @model_validator(mode='before') @classmethod - def _from_list(cls, data: Any) -> Any: + def _from_list(cls, data: object) -> object: if isinstance(data, list): if not 2 <= len(data) <= 3: msg = f'each link must be [expression, values] or [expression, values, sign], got {data!r}' @@ -561,7 +562,7 @@ def curve(self) -> tuple[PiecewiseLink, PiecewiseLink]: @field_validator('method', mode='wrap') @classmethod - def _check_method(cls, v: Any, handler: ValidatorFunctionWrapHandler) -> PiecewiseMethod: + def _check_method(cls, v: object, handler: ValidatorFunctionWrapHandler) -> PiecewiseMethod: try: return cast('PiecewiseMethod', handler(v)) except ValidationError: @@ -633,7 +634,7 @@ class SosBlock(_StrictBlock): @field_validator('type', mode='wrap') @classmethod - def _check_type(cls, v: Any, handler: ValidatorFunctionWrapHandler) -> SosType: + def _check_type(cls, v: object, handler: ValidatorFunctionWrapHandler) -> SosType: orders = ' or '.join(str(t) for t in sorted(SOS_TYPES)) msg = f'sos type must be {orders}, got {v!r}. A set of any other order is not a construct solvers carry.' if type(v) is not int: # True == 1 == 1.0, and a set of order True is nothing @@ -665,7 +666,7 @@ def undeclared_dimension(kind: str, name: str, dimension: str) -> str: return f"{kind} '{name}' references undeclared dimension '{dimension}'. Declare it under 'dimensions:'." -def _without_absence(value: Any) -> Any: +def _without_absence(value: object) -> object: """*value* with every absent entry stripped, recursively — see :meth:`Spec._drop_absence`.""" if not isinstance(value, dict): return value @@ -677,7 +678,7 @@ def _without_absence(value: Any) -> Any: return kept -def _is_absent(value: Any) -> bool: +def _is_absent(value: object) -> bool: """Whether *value* is a null or an infinite bound.""" if value is None: return True @@ -731,14 +732,32 @@ def relations_of(self, dimension: str) -> dict[str, RelationBlock]: @classmethod @override - def model_validate(cls, *args: Any, **kwargs: Any) -> Self: + def model_validate( + cls, + obj: object, + *, + strict: bool | None = None, + extra: ExtraValues | None = None, + from_attributes: bool | None = None, + context: object = None, + by_alias: bool | None = None, + by_name: bool | None = None, + ) -> Self: """Validate a mapping, raising this package's exception tree rather than pydantic's. ``__init__`` is not wrapped the same way, because defining one makes pydantic run every after-validator twice. """ try: - return super().model_validate(*args, **kwargs) + return super().model_validate( + obj, + strict=strict, + extra=extra, + from_attributes=from_attributes, + context=context, + by_alias=by_alias, + by_name=by_name, + ) except ValidationError as exc: raise schema_error(exc) from None @@ -758,16 +777,16 @@ def _check_version(cls, v: int) -> int: raise ValueError(msg) @model_serializer(mode='wrap') - def _drop_absence(self, handler: Any) -> dict[str, Any]: + def _drop_absence(self, handler: SerializerFunctionWrapHandler) -> dict[str, object]: """Absence is not serialised: a null, an infinite bound, a mapping that stripping emptied, a section declaring nothing. An empty list stays, being a value rather than an absence (``dims: []`` is a scalar). On the serializer so that ``model_dump``, :meth:`to_dict` and :meth:`to_yaml` agree. """ - return cast('dict[str, Any]', _without_absence(handler(self))) + return cast('dict[str, object]', _without_absence(handler(self))) - def to_dict(self) -> dict[str, Any]: + def to_dict(self) -> dict[str, object]: """The model as plain data. ``to_spec(m.to_dict())`` reproduces it.""" return self.model_dump() diff --git a/src/math_spec/piecewise.py b/src/math_spec/piecewise.py index 1df23845..bba5ae3b 100644 --- a/src/math_spec/piecewise.py +++ b/src/math_spec/piecewise.py @@ -13,7 +13,7 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING from math_spec._expression_parser import ComparisonNode from math_spec.degree import check_expression @@ -121,7 +121,7 @@ class _Block: or emitting a name the file already declares. """ - def __init__(self, schema: Spec, raw: dict[str, Any], name: str, pw: PiecewiseBlock) -> None: + def __init__(self, schema: Spec, raw: dict[str, object], name: str, pw: PiecewiseBlock) -> None: self.schema = schema self.raw = raw self.name = name @@ -143,9 +143,9 @@ def __init__(self, schema: Spec, raw: dict[str, Any], name: str, pw: PiecewiseBl self.ns = Namespace.of(schema) self.context = f"piecewise '{name}'" self.frame = self._validated_frame() - self.record: dict[str, Any] = {'block': raw['piecewise'][name], 'points': self.mask} + self.record: dict[str, object] = {'block': self._section('piecewise')[name], 'points': self.mask} - def expand(self) -> dict[str, Any]: + def expand(self) -> dict[str, object]: """Write the block's declarations, and return the record ``expanded_piecewise`` keeps for it.""" if self.nominated is not None: self._parameter( @@ -161,20 +161,30 @@ def expand(self) -> dict[str, Any]: # -- emitters ---------------------------------------------------------- + def _section(self, name: str) -> dict[str, object]: + """The *name* section of the raw model, created empty where the file declares none.""" + section = self.raw.setdefault(name, {}) + assert isinstance(section, dict), f'{name}: is a mapping in a validated model' + return section + + def _mask_dims(self, mask: str) -> list[str]: + """The dims of *mask*, the parameter masking the weights: the nominated values parameter's where it is derived from one.""" + return list(self.schema.parameters[self.nominated if self.nominated is not None else mask].dims) + def _parameter(self, name: str, dims: list[str], description: str) -> None: """A ``bool`` parameter the expansion derives.""" - self.raw.setdefault('parameters', {})[name] = {'dims': dims, 'dtype': 'bool', 'description': description} + self._section('parameters')[name] = {'dims': dims, 'dtype': 'bool', 'description': description} - def _weight(self, name: str, **fields: Any) -> None: + def _weight(self, name: str, **fields: object) -> None: """A variable over the frame and the breakpoint dim, masked as the block is.""" - self.raw['variables'][name] = { + self._section('variables')[name] = { 'dims': [*self.frame, self.pw.over], **({'where': self.mask} if self.mask else {}), **fields, } def _constraint(self, name: str, dims: list[str], expression: str, where: str | None = None) -> None: - self.raw['constraints'][name] = { + self._section('constraints')[name] = { 'dims': dims, **({'where': where} if where else {}), 'expression': expression, @@ -198,7 +208,7 @@ def _weights(self) -> None: f'({link.expression}) {link.sign} sum({self.lam} * {link.values}, over={d})', ) if self.pw.method == 'sos2': - self.raw.setdefault('sos', {})[self.name] = {'variable': self.lam, 'over': d, 'type': 2} + self._section('sos')[self.name] = {'variable': self.lam, 'over': d, 'type': 2} elif self.pw.method == 'adjacency': self._weight(self.seg, domain='binary', bounds={}) for suffix, where, rhs in gated: @@ -263,9 +273,7 @@ def _segment_lines(self) -> None: if mask: self.record['starts' if sense == '>=' else 'ends'] = at self._parameter( - at, - self.raw['parameters'][mask]['dims'], - f'the {"first" if sense == ">=" else "last"} breakpoint of each curve', + at, self._mask_dims(mask), f'the {"first" if sense == ">=" else "last"} breakpoint of each curve' ) self._constraint(cname, [*self.frame, d], f'({x_link.expression}) {sense} {x_link.values}', at) diff --git a/src/math_spec/typesetting/__init__.py b/src/math_spec/typesetting/__init__.py index 0575fe39..7bbdadd8 100644 --- a/src/math_spec/typesetting/__init__.py +++ b/src/math_spec/typesetting/__init__.py @@ -26,7 +26,7 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, Literal +from typing import TYPE_CHECKING, Literal, TypedDict, Unpack from math_spec.errors import SchemaError, did_you_mean from math_spec.piecewise import expand_piecewise @@ -66,10 +66,20 @@ } +class _Options(TypedDict, total=False): + """The keyword arguments :func:`typeset` takes, which the three per-format doors forward whole.""" + + symbols: str | Path | Mapping[str, object] | SymbolTable | None + standalone: bool + legend: bool + numbered: bool + inline_expressions: bool + + def _walk( - model: str | Path | dict[str, Any] | Spec, + model: str | Path | Mapping[str, object] | Spec, fmt: FormatName, - symbols: str | Path | Mapping[str, Any] | SymbolTable | None, + symbols: str | Path | Mapping[str, object] | SymbolTable | None, *, inline_expressions: bool, ) -> Walk: @@ -91,10 +101,10 @@ def _walk( def typeset( - model: str | Path | dict[str, Any] | Spec, + model: str | Path | Mapping[str, object] | Spec, fmt: FormatName, *, - symbols: str | Path | Mapping[str, Any] | SymbolTable | None = None, + symbols: str | Path | Mapping[str, object] | SymbolTable | None = None, standalone: bool = False, legend: bool = True, numbered: bool = True, @@ -148,11 +158,11 @@ def typeset( def typeset_declaration( - model: str | Path | dict[str, Any] | Spec, + model: str | Path | Mapping[str, object] | Spec, name: str, fmt: FormatName, *, - symbols: str | Path | Mapping[str, Any] | SymbolTable | None = None, + symbols: str | Path | Mapping[str, object] | SymbolTable | None = None, inline_expressions: bool = True, ) -> str: """Render one declaration as the bare line the document prints for it. @@ -199,16 +209,16 @@ def typeset_declaration( return walk.format.equation(walk.line(name)) -def to_latex(model: str | Path | dict[str, Any] | Spec, **options: Any) -> str: +def to_latex(model: str | Path | Mapping[str, object] | Spec, **options: Unpack[_Options]) -> str: """Render *model* as LaTeX (amsmath ``align``). See :func:`typeset`.""" return typeset(model, 'latex', **options) -def to_typst(model: str | Path | dict[str, Any] | Spec, **options: Any) -> str: +def to_typst(model: str | Path | Mapping[str, object] | Spec, **options: Unpack[_Options]) -> str: """Render *model* as Typst. See :func:`typeset`.""" return typeset(model, 'typst', **options) -def to_markdown(model: str | Path | dict[str, Any] | Spec, **options: Any) -> str: +def to_markdown(model: str | Path | Mapping[str, object] | Spec, **options: Unpack[_Options]) -> str: """Render *model* as GitHub-flavoured Markdown. See :func:`typeset`.""" return typeset(model, 'markdown', **options) diff --git a/src/math_spec/typesetting/symbols.py b/src/math_spec/typesetting/symbols.py index c6fb99d9..e7b1cd6e 100644 --- a/src/math_spec/typesetting/symbols.py +++ b/src/math_spec/typesetting/symbols.py @@ -13,7 +13,7 @@ from collections.abc import Mapping from dataclasses import dataclass, field from pathlib import Path -from typing import TYPE_CHECKING, Any, cast +from typing import TYPE_CHECKING, cast import math_spec.degree as degree from math_spec._yaml import read_yaml @@ -193,12 +193,13 @@ class SymbolTable: names: dict[str, str] = field(default_factory=dict) @classmethod - def load(cls, source: str | Path | Mapping[str, Any]) -> SymbolTable: + def load(cls, source: str | Path | Mapping[str, object]) -> SymbolTable: """A table from a YAML path or the mapping it parses to. Raises: - SchemaError: An unknown section, a malformed dimension, or a - ``notation:`` that is missing or not ``latex``/``typst``. + SchemaError: An unknown section, a section or a dimension that is + not a mapping, or a ``notation:`` that is missing or not + ``latex``/``typst``. """ raw = dict(source) if isinstance(source, Mapping) else read_yaml(Path(source)) unknown = set(raw) - {'notation', 'dimensions', 'names'} @@ -215,7 +216,7 @@ def load(cls, source: str | Path | Mapping[str, Any]) -> SymbolTable: indices: dict[str, str] = {} sets: dict[str, str] = {} - for dim, spec in (raw.get('dimensions') or {}).items(): + for dim, spec in _section(raw, 'dimensions').items(): if not isinstance(spec, Mapping): msg = f"symbol table: dimension '{dim}' must be a mapping like {{index: t, set: '\\\\mathcal{{T}}'}}" raise SchemaError(msg) @@ -232,7 +233,7 @@ def load(cls, source: str | Path | Mapping[str, Any]) -> SymbolTable: notation=cast('Notation', notation), indices=indices, sets=sets, - names={k: str(v) for k, v in (raw.get('names') or {}).items()}, + names={k: str(v) for k, v in _section(raw, 'names').items()}, ) def checked_against(self, schema: _ExpandedSpec) -> SymbolTable: @@ -250,5 +251,20 @@ def checked_against(self, schema: _ExpandedSpec) -> SymbolTable: return self +def _section(raw: Mapping[str, object], name: str) -> Mapping[str, object]: + """The *name* section of a symbol table as the mapping it has to be, empty where it is absent or null. + + Raises: + SchemaError: The section is something else, such as a list. + """ + section = raw.get(name) + if section is None: + return {} + if not isinstance(section, Mapping): + msg = f'symbol table: {name}: must be a mapping of names to entries, got {type(section).__name__}.' + raise SchemaError(msg) + return cast('Mapping[str, object]', section) + + def _unknown_entry(name: str, section: str, known: set[str]) -> str: return f"symbol table: '{name}' under {section}: is not declared by the model. {did_you_mean(name, known)}" diff --git a/src/math_spec/validation.py b/src/math_spec/validation.py index 5169a150..e292840e 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -6,7 +6,8 @@ from __future__ import annotations -from typing import TYPE_CHECKING, Any, Literal, assert_never, overload +from collections.abc import Mapping +from typing import TYPE_CHECKING, Literal, assert_never, overload import math_spec.degree as degree from math_spec._expression_parser import ( @@ -55,7 +56,7 @@ from math_spec.program import WhereNode -def to_spec(model: str | Path | dict[str, Any] | Spec) -> Spec: +def to_spec(model: str | Path | Mapping[str, object] | Spec) -> Spec: """Load and validate a model definition — the language's front door. Everything decidable without data is decided here: schema shape, every @@ -80,7 +81,7 @@ def to_spec(model: str | Path | dict[str, Any] | Spec) -> Spec: raise SchemaError(msg) if isinstance(model, Spec): return model - return Spec.model_validate(model if isinstance(model, dict) else read_model(model)) + return Spec.model_validate(model if isinstance(model, Mapping) else read_model(model)) def _once(errors: list[str]) -> str: diff --git a/tests/typesetting/test_symbols.py b/tests/typesetting/test_symbols.py index cabd21d4..b0ea7db1 100644 --- a/tests/typesetting/test_symbols.py +++ b/tests/typesetting/test_symbols.py @@ -109,6 +109,8 @@ def test_a_named_expression_has_a_legend_row_exactly_while_its_symbol_prints(nam id='a-table-still-carrying-descriptions', ), pytest.param({'dimensions': {'generator': {'letter': 'g'}}}, 'unknown key', id='an-unknown-key'), + pytest.param({'dimensions': ['generator']}, 'dimensions: must be a mapping', id='a-section-that-is-a-list'), + pytest.param({'names': 'p_max'}, 'names: must be a mapping', id='a-section-that-is-a-string'), ], ) def test_an_entry_naming_nothing_is_an_error_with_the_near_miss(symbols, match): From e5f73aa795aa475d6c0f629f2cd80b437fa01d01 Mon Sep 17 00:00:00 2001 From: Claude Date: Fri, 18 Sep 2026 23:13:27 +0000 Subject: [PATCH 2/2] refactor: the ordering helper names the three kinds a literal comes in rather than a protocol A constrained type variable says what the four-method protocol said, in one line. The two type statements are the plain unions every other alias in the package is, and the symbol table section needs no cast once narrowed. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_01EEeM2YoAk4Xr2uWB5qwMsH --- src/math_spec/_where_parser.py | 5 ++--- src/math_spec/exclusivity.py | 16 ++++------------ src/math_spec/typesetting/symbols.py | 2 +- 3 files changed, 7 insertions(+), 16 deletions(-) diff --git a/src/math_spec/_where_parser.py b/src/math_spec/_where_parser.py index 4e6b805d..ed9b0123 100644 --- a/src/math_spec/_where_parser.py +++ b/src/math_spec/_where_parser.py @@ -28,14 +28,13 @@ NotNode, OrNode, PredicateOperator, + WhereNode, where_children, ) if TYPE_CHECKING: from collections.abc import Callable - from math_spec.program import WhereNode - # --------------------------------------------------------------------------- # AST nodes # --------------------------------------------------------------------------- @@ -94,7 +93,7 @@ class UnresolvedComparisonNode: #: Every node a parsed where string is built of: the connectives and literals, #: the unresolved leaves, and the arithmetic and the two side nodes under a #: comparison. What the depth measurement walks. -type _ParsedWhere = WhereNode | UnresolvedWhereNode | ArithmeticNode | ColumnNode | QuotedNode +_ParsedWhere = WhereNode | UnresolvedWhereNode | ArithmeticNode | ColumnNode | QuotedNode # --------------------------------------------------------------------------- diff --git a/src/math_spec/exclusivity.py b/src/math_spec/exclusivity.py index 8e676743..a91e7fae 100644 --- a/src/math_spec/exclusivity.py +++ b/src/math_spec/exclusivity.py @@ -19,7 +19,7 @@ import math from dataclasses import dataclass from enum import Enum -from typing import TYPE_CHECKING, Literal, Protocol, Self, assert_never +from typing import TYPE_CHECKING, Literal, assert_never from math_spec.program import ( AndNode, @@ -130,16 +130,7 @@ class Special(Enum): #: What a where comparison is written against: a number, a label, or a date. #: ``position()`` counts in integers, which are numbers here. -type _Literal = float | str | datetime.date - - -class _Ordered(Protocol): - """A value that orders against its own kind — what one atom's truth in one cell asks of both sides.""" - - def __lt__(self, other: Self, /) -> bool: ... - def __le__(self, other: Self, /) -> bool: ... - def __gt__(self, other: Self, /) -> bool: ... - def __ge__(self, other: Self, /) -> bool: ... +_Literal = float | str | datetime.date @dataclass(frozen=True) @@ -484,7 +475,8 @@ def _compare(value: Cell, op: PredicateOperator, literal: _Literal) -> bool: raise AssertionError(msg) -def _ordered[T: _Ordered](left: T, op: PredicateOperator, right: T) -> bool: +def _ordered[T: (float, str, datetime.date)](left: T, op: PredicateOperator, right: T) -> bool: + """One comparison between two values of one kind — the kinds a literal comes in.""" match op: case '==': return left == right diff --git a/src/math_spec/typesetting/symbols.py b/src/math_spec/typesetting/symbols.py index e7b1cd6e..995280cb 100644 --- a/src/math_spec/typesetting/symbols.py +++ b/src/math_spec/typesetting/symbols.py @@ -263,7 +263,7 @@ def _section(raw: Mapping[str, object], name: str) -> Mapping[str, object]: if not isinstance(section, Mapping): msg = f'symbol table: {name}: must be a mapping of names to entries, got {type(section).__name__}.' raise SchemaError(msg) - return cast('Mapping[str, object]', section) + return section def _unknown_entry(name: str, section: str, known: set[str]) -> str: