Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 38 additions & 11 deletions src/math_spec/resolution.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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 -------------------------------------------------------

Expand All @@ -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)
Expand Down Expand Up @@ -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:
Expand Down Expand Up @@ -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()
Expand All @@ -519,11 +532,15 @@ 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
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)

Expand All @@ -536,14 +553,16 @@ 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 — '
f'dual(<constraint>). Name the constraint whose row dual you want.'
)
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)

Expand Down Expand Up @@ -576,18 +595,21 @@ def _relation_ref(
)
return value
name = names[0]

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
named = {k: r for k, r in read.items() if r is not None}
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)):
Expand Down Expand Up @@ -744,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: <the columns a row is identified by>, "
f'values: <the columns they determine>}}.'
Expand Down Expand Up @@ -1257,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'
Expand Down
89 changes: 2 additions & 87 deletions src/math_spec/validation.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,37 +7,24 @@
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
from math_spec.errors import LanguageError, SchemaError, prefixed
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 (
Expand All @@ -46,7 +33,6 @@
ResolvedAssumption,
ResolvedConstraint,
mask_of,
names_in,
resolve_expression,
resolve_where_text,
)
Expand Down Expand Up @@ -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():
Expand Down Expand Up @@ -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)
72 changes: 68 additions & 4 deletions tests/test_expansion.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'],
Expand Down Expand Up @@ -223,19 +223,83 @@ 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',
),
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)


@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'),
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 `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'}},
)
Loading