From 3d7584dad3636fa4dc5be7598cde84ebc4845f5f Mon Sep 17 00:00:00 2001 From: Claude Date: Tue, 22 Sep 2026 20:04:07 +0000 Subject: [PATCH 1/2] fix(language): a macro template nothing calls is held to every rule a call site is The template check was a second walker that looked at names only. It is now the resolver, with the template's formals left bare, so a label parameter used as a value or an unknown relation column is refused at load. Co-Authored-By: Claude Opus 5.5 Claude-Session: https://claude.ai/code/session_017XwKgY5wZXv1bKgkfCW2q1 --- src/math_spec/resolution.py | 29 +++++++++--- src/math_spec/validation.py | 89 +------------------------------------ tests/test_expansion.py | 35 +++++++++++++-- 3 files changed, 58 insertions(+), 95 deletions(-) diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index 7f0d5c6f..cfddcb15 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -288,16 +288,21 @@ def resolve_expression( ns: Namespace, context: str, errors: list[str], + *, + formals: frozenset[str] = frozenset(), ) -> ParsedNode | None: """Rewrite every ``NameNode`` under *node* to a typed node, checking operator call shapes on the way. + A name in *formals* stays bare, so a macro template is checked by the + rules a call site is, before anything calls it. + Returns: The typed tree, or ``None`` once anything failed — appending to *errors* rather than raising, so a caller collecting problems across a whole schema reports them together. """ before = len(errors) - resolved = _Resolver(ns, context, errors).expression(node) + resolved = _Resolver(ns, context, errors, formals=formals).expression(node) return None if len(errors) > before else resolved @@ -350,13 +355,19 @@ class _Resolver: appended to ``errors``; the public doors discard the tree once ``errors`` grew, which is what lets a connective's children be typed as resolved. ``self_variable`` is the variable whose own ``where`` is being read, which - may not ask whether it exists. + may not ask whether it exists. ``formals`` are a macro template's formals, + which stay bare: a formal has no kind until a call site binds it. """ ns: Namespace context: str errors: list[str] self_variable: str | None = None + formals: frozenset[str] = frozenset() + + def _formal(self, value: ArithmeticNode) -> bool: + """Whether *value* is a formal, left for the call site to bind.""" + return isinstance(value, NameNode) and value.name in self.formals # -- expressions ------------------------------------------------------- @@ -374,7 +385,7 @@ def _arith(self, node: ArithmeticNode, *, amount: bool = False) -> ArithmeticNod numeric check here stands aside for it. A quoted keyword or a name list in arithmetic arrives through a macro formal bound to one. """ - if isinstance(node, NumberNode | VariableNode | ParameterNode | DualNode | KwargNode): + if isinstance(node, NumberNode | VariableNode | ParameterNode | DualNode | KwargNode) or self._formal(node): return node if isinstance(node, NameNode): return self._name(node, amount=amount) @@ -429,7 +440,7 @@ def _name(self, node: NameNode, *, amount: bool) -> ArithmeticNode: ) return node case _: - self.errors.append(self.ns.unknown(node.name, self.context, allow_dims=False)) + self.errors.append(self.ns.unknown(node.name, self.context, allow_dims=False, formals=self.formals)) return node def _call(self, node: FunctionCallNode) -> ArithmeticNode: @@ -495,6 +506,8 @@ def _amount(self, value: ArithmeticNode, operator: str, key: str) -> ArithmeticN def _edge(self, value: ArithmeticNode, operator: str) -> ArithmeticNode: """``edge=``: the closed keyword ``wrap``, or a number to contribute; a name here is a typo.""" + if self._formal(value): + return value if isinstance(value, KeywordNode): if value.value == EDGE_WRAP: return EdgeNode() @@ -519,6 +532,8 @@ def _edge(self, value: ArithmeticNode, operator: str) -> ArithmeticNode: def _dim_ref(self, value: ArithmeticNode, operator: str, key: str) -> ArithmeticNode: """An operator kwarg whose *value* must name a declared dimension.""" + if self._formal(value): + return value if not isinstance(value, NameNode): self.errors.append(f'{self.context}: {operator}({key}=...) must name a dimension.') return value @@ -536,6 +551,8 @@ def _dual(self, node: FunctionCallNode) -> ArithmeticNode: (:mod:`math_spec.validation`); this pass only types the name. """ (value,) = node.args + if self._formal(value): + return node if not isinstance(value, NameNode): self.errors.append( f'{self.context}: dual() takes the name of a declared constraint, written bare — ' @@ -543,7 +560,7 @@ def _dual(self, node: FunctionCallNode) -> ArithmeticNode: ) return node if value.name not in self.ns.constraints: - self.errors.append(self.ns.unknown_constraint(value.name, self.context)) + self.errors.append(self.ns.unknown_constraint(value.name, self.context, formals=self.formals)) return node return DualNode(value.name) @@ -576,6 +593,8 @@ def _relation_ref( ) return value name = names[0] + if name in self.formals or any(n in self.formals for v in roles.values() for n in names_in(v)): + return value if (problem := self._not_a_relation(name, operator, key)) is not None: self.errors.append(problem) diff --git a/src/math_spec/validation.py b/src/math_spec/validation.py index e976ebce..f62df8d4 100644 --- a/src/math_spec/validation.py +++ b/src/math_spec/validation.py @@ -7,29 +7,17 @@ from __future__ import annotations from collections.abc import Mapping -from typing import TYPE_CHECKING, Literal, assert_never, overload +from typing import TYPE_CHECKING, Literal, overload import math_spec.degree as degree from math_spec._expression_parser import ( ArithmeticNode, - BinaryOperatorNode, CaseArm, CasesNode, ComparisonNode, DefinitionNode, - DualNode, - FunctionCallNode, - KeywordNode, - KwargNode, - NameListNode, - NameNode, - NumberNode, - ParameterNode, ParsedNode, - UnaryOperatorNode, - VariableNode, case_context, - children, ) from math_spec._yaml import read_model from math_spec.dimensions import check_schema @@ -37,7 +25,6 @@ from math_spec.exclusivity import overlapping from math_spec.expansion import expand, parse_and_expand, parse_template from math_spec.model import AssumptionBlock, Spec -from math_spec.operators import BUILTINS, call_shape_error, unknown_operator_message from math_spec.piecewise import assumptions_of from math_spec.program import BooleanLiteral, Mask, VariableDefined from math_spec.resolution import ( @@ -46,7 +33,6 @@ ResolvedAssumption, ResolvedConstraint, mask_of, - names_in, resolve_expression, resolve_where_text, ) @@ -141,7 +127,7 @@ def validate_expressions(schema: Spec) -> Resolved: f'ambiguous with the dimension itself.' for f in sorted(formals & ns.dimensions) ) - _check_template_names(body_ast, context, ns, formals, errors) + resolve_expression(body_ast, ns, context, errors, formals=formals) expressions: dict[str, CasesNode | DefinitionNode] = {} for ename, block in schema.expressions.items(): @@ -373,74 +359,3 @@ def _check_expression( ) return None return resolved - - -def _check_template_names( - node: ArithmeticNode, - context: str, - ns: Namespace, - formals: frozenset[str], - errors: list[str], -) -> None: - """Check a macro body's names and call shapes, treating formals as bound — not resolution, since a formal has no kind until a call site binds it. - - An operator call is refused by its signature here, as at a call site, so a - keyword the operator does not declare is caught in a template nothing calls. - A case arm's value only: its ``when`` is the declaration's, checked there. - """ - if isinstance(node, NumberNode | VariableNode | ParameterNode | DualNode | KwargNode | KeywordNode | NameListNode): - return - - if isinstance(node, NameNode): - if node.name not in formals and ns.kind(node.name) is None: - errors.append(ns.unknown(node.name, context, allow_dims=False, formals=formals)) - return - - if isinstance(node, UnaryOperatorNode | BinaryOperatorNode | CasesNode | DefinitionNode): - for child in children(node): - _check_template_names(child, context, ns, formals, errors) - return - - if isinstance(node, FunctionCallNode): - builtin = BUILTINS.get(node.name) - if builtin is None: - errors.append(f'{context}: {unknown_operator_message(node.name)}') - else: - shape_error = call_shape_error(node.name, len(node.args), node.kwargs) - if shape_error is not None: - errors.append(f'{context}: {shape_error}') - if node.name == 'dual': - errors.extend( - ns.unknown_constraint(arg.name, context, formals=formals) - for arg in node.args - if isinstance(arg, NameNode) and arg.name not in formals and arg.name not in ns.constraints - ) - return - for arg in node.args: - _check_template_names(arg, context, ns, formals, errors) - for kwarg, value in node.kwargs.items(): - with_relation = builtin is not None and any(k in node.kwargs for k in builtin.relation_kwargs) - match builtin.kind_of(kwarg, with_relation=with_relation) if builtin else 'value': - case 'dimension': - if isinstance(value, NameNode) and value.name not in ns.dimensions | formals: - errors.append( - f'{context}: {node.name}({kwarg}={value.name}) does not name a ' - f'declared dimension or a formal of this macro.' - ) - case 'relation': - errors.extend( - f'{context}: {node.name}({kwarg}={one}) does not name a relation or a formal of this macro.' - for one in names_in(value) - if one not in formals and ns.kind(one) != 'relation' - ) - case 'value': - _check_template_names(value, context, ns, formals, errors) - case 'role': - pass - case 'edge': - pass # a keyword or a number: nothing in it to name - case None: - pass # a keyword the operator does not declare; the shape error above named it - return - - assert_never(node) diff --git a/tests/test_expansion.py b/tests/test_expansion.py index b79bc6e6..f1eb067f 100644 --- a/tests/test_expansion.py +++ b/tests/test_expansion.py @@ -13,7 +13,7 @@ from math_spec._expression_parser import ComparisonNode, DefinitionNode, parse_expression, with_children from math_spec.errors import LanguageError from math_spec.expansion import parse_and_expand -from tests.fixtures import DISPATCH_MODEL, schema_of +from tests.fixtures import DISPATCH_MODEL, SMALL_MODEL, schema_of WEIGHTED_SUM = { 'args': ['array', 'weights'], @@ -223,8 +223,8 @@ def test_macro_collisions_rejected(patch, match): ), pytest.param( {'grouped': {'args': ['x'], 'template': 'sum(x, by=[nope, also])'}}, - r"Macro 'grouped'.*sum\(by=nope\) does not name a relation", - id='a-typo-in-a-relation-list', + r"Macro 'grouped'.*sum\(by=\[nope, also\]\) names 2 relations", + id='a-list-of-relations', ), ], ) @@ -234,8 +234,37 @@ def test_macro_templates_validated_even_when_unused(macros, match): schema(macros=macros) +@pytest.mark.parametrize( + ('template', 'match'), + [ + pytest.param('x * tag', "Macro 'm': 'tag' is declared dtype: str", id='a-label-parameter-as-a-value'), + pytest.param('sum(x, by=lk, over=nope, into=h)', "over=nope names no column of 'lk'", id='a-typo-in-a-column'), + ], +) +def test_a_template_is_held_to_the_rules_a_call_site_is(template, match): + """A template nothing calls was checked for names only: a label parameter or an unknown column passed load.""" + with pytest.raises(LanguageError, match=match): + schema_of(SMALL_MODEL, macros={'m': {'args': ['x'], 'template': template}}) + + @pytest.mark.parametrize('fragment', ['my_python_helper', 'macros:', 'escape']) def test_an_unknown_operator_is_refused_at_load_with_the_rewrite(fragment): with pytest.raises(LanguageError) as exc: schema(constraints={'c': {'dims': ['snapshot'], 'expression': 'my_python_helper(p) <= load'}}) assert fragment in str(exc.value) + + +@pytest.mark.parametrize( + ('formals', 'template'), + [ + pytest.param(['x', 'e'], 'shift(x, along=g, offset=1, edge=e)', id='an-edge'), + pytest.param(['row'], 'dual(row)', id='a-constraint'), + pytest.param(['x', 'rel', 'a', 'b'], 'sum(x, by=rel, over=a, into=b)', id='a-relation-and-its-columns'), + pytest.param(['x', 'a', 'b'], 'sum(x, by=lk, over=a, into=b)', id='the-columns-of-a-declared-relation'), + ], +) +def test_a_formal_stands_where_a_call_site_will_bind_it(formals, template): + """A formal has no kind until a call binds it, so the template check leaves it bare in every slot.""" + assert ( + schema_of(SMALL_MODEL, macros={'m': {'args': formals, 'template': template}}).macros['m'].template == template + ) From 89335fc8cb667b07c1cefd74d6b3b825363830d7 Mon Sep 17 00:00:00 2001 From: Claude Date: Wed, 23 Sep 2026 11:23:06 +0000 Subject: [PATCH 2/2] fix(language): a template's formal along= builds no partition, and a typo in by= is refused beside formal columns A formal along= beside a by= was handed to the partition as a dimension and refused as one the relation has no key column over, so a valid template was refused. A relation that names nothing loaded once the columns beside it were formals, because the formals sent the call back first. The over= and by= refusals inside a template say "or a formal of this macro" again. Co-Authored-By: Claude Fable 5.1 Claude-Session: https://claude.ai/code/session_015h57WkBDnpxrknuJ5zZy9F --- src/math_spec/resolution.py | 22 ++++++++++++++------- tests/test_expansion.py | 39 +++++++++++++++++++++++++++++++++++-- 2 files changed, 52 insertions(+), 9 deletions(-) diff --git a/src/math_spec/resolution.py b/src/math_spec/resolution.py index cfddcb15..9b1d0e38 100644 --- a/src/math_spec/resolution.py +++ b/src/math_spec/resolution.py @@ -538,7 +538,9 @@ def _dim_ref(self, value: ArithmeticNode, operator: str, key: str) -> Arithmetic self.errors.append(f'{self.context}: {operator}({key}=...) must name a dimension.') return value if value.name not in self.ns.dimensions: - self.errors.append(_undeclared_dim(self.context, operator, f'{key}={value.name}', value.name, self.ns)) + self.errors.append( + _undeclared_dim(self.context, operator, f'{key}={value.name}', value.name, self.ns, self.formals) + ) return value return DimensionNode(value.name) @@ -593,12 +595,13 @@ def _relation_ref( ) return value name = names[0] - if name in self.formals or any(n in self.formals for v in roles.values() for n in names_in(v)): + if name in self.formals: return value - if (problem := self._not_a_relation(name, operator, key)) is not None: self.errors.append(problem) return value + if any(n in self.formals for v in roles.values() for n in names_in(v)): + return value read = {k: self._role_name(v, operator, k) for k, v in roles.items()} if any(r is None for r in read.values()): return value @@ -606,7 +609,7 @@ def _relation_ref( if operator in ('shift', 'sum_back'): if 'within' not in named: return value # the call shape refused it already, with the wording that names the rewrite - over_dim = over.name if isinstance(over, NameNode | DimensionNode) else None + over_dim = over.name if isinstance(over, NameNode | DimensionNode) and not self._formal(over) else None partition = self._partition(name, operator, over_dim, named['within']) return value if partition is None else PartitionNode(partition) if not ({'over', 'into'} <= set(named)): @@ -763,7 +766,7 @@ def _not_a_relation(self, name: str, operator: str, key: str) -> str | None: f'{key}= takes a relation — the named map out of a dimension.\n{hint}' ) return ( - f'{context}: {operator}({key}={name}) does not name a relation. ' + f'{context}: {operator}({key}={name}) does not name a relation{_or_a_formal(self.formals)}. ' f'{did_you_mean(name, ns.relations, label="Relations")}\n' f"Declare it under 'relations:' — {name}: {{key: , " f'values: }}.' @@ -1276,15 +1279,20 @@ def _not_a_number(name: str, dtype: str, context: str) -> str: ) -def _undeclared_dim(context: str, operator: str, call: str, name: str, ns: Namespace) -> str: +def _undeclared_dim(context: str, operator: str, call: str, name: str, ns: Namespace, formals: frozenset[str]) -> str: return ( - f'{context}: {operator}({call}) does not name a declared dimension. ' + f'{context}: {operator}({call}) does not name a declared dimension{_or_a_formal(formals)}. ' f'{did_you_mean(name, ns.dimensions, label="Dimensions")}\n' f"Declare '{name}' under 'dimensions:', or fix the typo — an unknown " f'dimension makes {operator}() a silent no-op rather than an error.' ) +def _or_a_formal(formals: frozenset[str]) -> str: + """The words a refusal inside a template adds, since a formal would have stood there too.""" + return ' or a formal of this macro' if formals else '' + + def _declared_as(ns: Namespace, name: str) -> str: kind = ns.kind(name) return f'a {kind}' if kind else 'not declared' diff --git a/tests/test_expansion.py b/tests/test_expansion.py index f1eb067f..573e034d 100644 --- a/tests/test_expansion.py +++ b/tests/test_expansion.py @@ -226,10 +226,25 @@ def test_macro_collisions_rejected(patch, match): r"Macro 'grouped'.*sum\(by=\[nope, also\]\) names 2 relations", id='a-list-of-relations', ), + pytest.param( + {'grouped': {'args': ['x', 'a', 'b'], 'template': 'sum(x, by=nope, over=a, into=b)'}}, + r"Macro 'grouped'.*sum\(by=nope\) does not name a relation or a formal of this macro", + id='a-typo-in-a-relation-beside-formal-columns', + ), + pytest.param( + {'reduced': {'args': ['x'], 'template': 'sum(x, over=nope)'}}, + r"Macro 'reduced'.*sum\(over=nope\) does not name a declared dimension or a formal of this macro", + id='a-typo-in-a-dimension', + ), ], ) def test_macro_templates_validated_even_when_unused(macros, match): - """A typo in a template the model never calls is still caught at load.""" + """A typo in a template the model never calls is still caught at load. + + The relation beside formal columns loaded once the columns were formals, + because the formals sent the call back before the relation's name was + read. + """ with pytest.raises(LanguageError, match=match): schema(macros=macros) @@ -261,10 +276,30 @@ def test_an_unknown_operator_is_refused_at_load_with_the_rewrite(fragment): pytest.param(['row'], 'dual(row)', id='a-constraint'), pytest.param(['x', 'rel', 'a', 'b'], 'sum(x, by=rel, over=a, into=b)', id='a-relation-and-its-columns'), pytest.param(['x', 'a', 'b'], 'sum(x, by=lk, over=a, into=b)', id='the-columns-of-a-declared-relation'), + pytest.param( + ['x', 'd'], 'shift(x, along=d, offset=1, by=lk, within=h)', id='the-dimension-a-partition-steps-along' + ), + pytest.param( + ['x', 'd'], 'sum_back(x, along=d, window=2, by=lk, within=h)', id='the-dimension-a-window-runs-along' + ), ], ) def test_a_formal_stands_where_a_call_site_will_bind_it(formals, template): - """A formal has no kind until a call binds it, so the template check leaves it bare in every slot.""" + """A formal has no kind until a call binds it, so the template check leaves it bare in every slot. + + A formal `along=` beside a `by=` was handed to the partition as if it were + a dimension, and refused as one the relation has no key column over. + """ assert ( schema_of(SMALL_MODEL, macros={'m': {'args': formals, 'template': template}}).macros['m'].template == template ) + + +def test_a_call_binding_the_dimension_a_partition_steps_along_loads(): + """The call site is where the formal gets its kind, so the partition is read there.""" + template = 'shift(x, along=d, offset=1, by=lk, within=h)' + schema_of( + SMALL_MODEL, + macros={'m': {'args': ['x', 'd'], 'template': template}}, + constraints={'c': {'dims': ['g'], 'expression': 'm(p, g) <= 1'}}, + )