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
15 changes: 11 additions & 4 deletions src/math_spec/_expression_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -412,8 +412,12 @@ def with_children(node: ArithmeticNode, recurse: Callable[[ArithmeticNode], Arit
# ---------------------------------------------------------------------------


def _build_grammar() -> pp.ParserElement:
"""``inf`` is a ``pp.Keyword`` rather than a ``pp.Literal``, which would match the prefix of ``inflow``."""
def _build_grammar() -> tuple[pp.ParserElement, pp.ParserElement]:
"""The arithmetic grammar, and the expression grammar that puts one comparison over it.

``inf`` is a ``pp.Keyword`` rather than a ``pp.Literal``, which would
match the prefix of ``inflow``.
"""
arith = pp.Forward()

inf_literal = (pp.Keyword('.inf') | pp.Keyword('inf')).set_parse_action(lambda: NumberNode(float('inf')))
Expand Down Expand Up @@ -449,9 +453,10 @@ def _build_grammar() -> pp.ParserElement:
arith <<= add_sub

comparator = pp.one_of(list(get_args(ComparisonOperator)))
return (arith + pp.Optional(comparator + arith)).set_parse_action(
expression = (arith + pp.Optional(comparator + arith)).set_parse_action(
lambda t: ComparisonNode(t[1], t[0], t[2]) if len(t) == 3 else t[0]
)
return arith, expression


def _make_func_call(tokens: pp.ParseResults) -> FunctionCallNode:
Expand Down Expand Up @@ -485,7 +490,9 @@ def _make_power(tokens: pp.ParseResults) -> ArithmeticNode:
return items[0] if len(items) == 1 else BinaryOperatorNode('**', items[0], items[2])


_GRAMMAR = _build_grammar()
#: The arithmetic half on its own, for the where grammar to put a predicate's
#: comparator over: one grammar for what a side may say, wherever it stands.
ARITHMETIC, _GRAMMAR = _build_grammar()


#: How deep a tree the language admits. Every pass over an expression recurses,
Expand Down
151 changes: 77 additions & 74 deletions src/math_spec/_where_parser.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,12 @@

"""The where-string grammar and the ``Unresolved*`` nodes it emits, package-private.

The resolved vocabulary lives in :mod:`math_spec.program`.
A where string is a boolean algebra over comparisons, and a comparison's
sides are the expression grammar's own arithmetic. What a side *is* — a
parameter, a dimension, a relation column, a ``position()`` — only the schema
knows, so the grammar hands both sides over bare and
:mod:`math_spec.resolution` reads them. The resolved vocabulary lives in
:mod:`math_spec.program`.
"""

from __future__ import annotations
Expand All @@ -15,14 +20,21 @@

import pyparsing as pp

from math_spec._expression_parser import NAME, REAL, 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,
WhereNode,
where_children,
)

if TYPE_CHECKING:
from collections.abc import Callable

from math_spec.program import WhereNode

# ---------------------------------------------------------------------------
# AST nodes
# ---------------------------------------------------------------------------
Expand All @@ -36,109 +48,88 @@ class UnresolvedNameNode:


@dataclass(frozen=True)
class UnresolvedComparisonNode:
"""A comparison against an unresolved name. ``resolution.py`` types it."""
class ColumnNode:
"""``relation.column`` on a side of a comparison — the one place the language names a column."""

name: str
op: PredicateOperator
value: float | str
#: Whether the right-hand side arrived in quotes. A bare word is ambiguous
#: — it may name a declaration — and resolution refuses it for that reason;
#: a quoted one is unambiguously a label, which is the only way to write
#: ``combined-cycle`` or a date. Consumed by resolution, never lowered.
quoted: bool = False
relation: str
column: str

@property
def shown(self) -> str:
"""The column as the file wrote it, for an error message."""
return f'{self.relation}.{self.column}'

@dataclass(frozen=True)
class UnresolvedPositionNode:
"""``position(dim[, by=relation, within=columns]) <op> i`` before the names are checked; ``resolution.py`` types it."""

dimension: str
op: PredicateOperator
position: int
by: str | None = None
into: tuple[str, ...] | None = None
@dataclass(frozen=True)
class QuotedNode:
"""A right-hand side that arrived in quotes.

A bare word is ambiguous — it may name a declaration — and resolution
refuses it for that reason; a quoted one is unambiguously a label, which
is the only way to write ``combined-cycle`` or a date.
"""

#: What resolution rewrites away on the where side — the three nodes whose
#: left-hand side is still a name the schema has not been asked about.
UnresolvedWhereNode = UnresolvedNameNode | UnresolvedComparisonNode | UnresolvedPositionNode
value: str


# ---------------------------------------------------------------------------
# Grammar
# ---------------------------------------------------------------------------
@dataclass(frozen=True)
class UnresolvedComparisonNode:
"""``side <op> side`` before the sides are read; ``resolution.py`` decides what each is.

A side is the expression grammar's arithmetic, so a name, a number and a
``position(...)`` call all arrive as the nodes an expression would carry
them in; a relation column and a quoted label have nodes of their own.
"""

class _Quoted(str):
"""A right-hand side that arrived in quotes; :func:`_comparison` turns it back into a flag."""
left: ArithmeticNode | ColumnNode
op: PredicateOperator
right: ArithmeticNode | ColumnNode | QuotedNode

__slots__ = ()

#: What resolution rewrites away on the where side — the two nodes whose
#: leaves are still names the schema has not been asked about.
UnresolvedWhereNode = UnresolvedNameNode | UnresolvedComparisonNode

def _position_comparison(tokens: pp.ParseResults) -> UnresolvedPositionNode:
"""``position(dim[, by=relation, within=columns]) <op> i`` off the tokens the grammar captured."""
dimension, *call, op, at = tokens
by = str(call[0]) if call else None
into = tuple(str(token) for token in call[1]) if len(call) > 1 else None
return UnresolvedPositionNode(str(dimension), op, at, by, into)
#: 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.
_ParsedWhere = WhereNode | UnresolvedWhereNode | ArithmeticNode | ColumnNode | QuotedNode


def _comparison(tokens: pp.ParseResults) -> UnresolvedComparisonNode:
"""``name <op> literal`` off the tokens the grammar captured, the quoted marker turned into a flag."""
name, op, value = tokens
quoted = isinstance(value, _Quoted)
return UnresolvedComparisonNode(str(name), op, str(value) if quoted else value, quoted)
# ---------------------------------------------------------------------------
# Grammar
# ---------------------------------------------------------------------------


def _build_where_grammar() -> pp.ParserElement:
"""Build the pyparsing grammar for where strings.

Both quote characters are accepted because YAML already owns one of them.
``NOT`` binds tightest, then ``AND``, then ``OR``. ``position(...)`` leads
the alternation, since ``position`` would otherwise be read as a bare name.
``NOT`` binds tightest, then ``AND``, then ``OR``. A comparison is tried
before a bare name, since its left side begins with one.
"""
where_expr = pp.Forward()

true_lit = pp.CaselessKeyword('True').set_parse_action(lambda: BooleanLiteralNode(True))
false_lit = pp.CaselessKeyword('False').set_parse_action(lambda: BooleanLiteralNode(False))

# pyrefly: ignore[implicit-any-lambda]
number = pp.Regex(rf'-?({REAL}|\d+)').set_parse_action(lambda t: float(t[0]))
# pyrefly: ignore[implicit-any-lambda]
position = pp.Regex(r'-?\d+').set_parse_action(lambda t: int(t[0]))

name = pp.Regex(NAME)

# pyrefly: ignore[implicit-any-lambda]
column = pp.Regex(rf'({NAME})\.({NAME})').set_parse_action(lambda t: ColumnNode(*t[0].split('.')))
quoted = (pp.QuotedString("'", esc_char='\\') | pp.QuotedString('"', esc_char='\\')).set_parse_action(
lambda t: _Quoted(t[0])
)

column = pp.Regex(rf'{NAME}(\.{NAME})?')
columns = name | (pp.Suppress('[') + pp.DelimitedList(name) + pp.Suppress(']'))
grouped_within = pp.Group(pp.Suppress(',') + pp.Suppress(pp.Keyword('within')) + pp.Suppress('=') + columns)
grouped_by = (
pp.Suppress(',') + pp.Suppress(pp.Keyword('by')) + pp.Suppress('=') + name + pp.Optional(grouped_within)
# pyrefly: ignore[implicit-any-lambda]
lambda t: QuotedNode(t[0])
)
comparator = pp.one_of(list(get_args(PredicateOperator)))

position_call = (
pp.Suppress(pp.Keyword('position')) + pp.Suppress('(') + name + pp.Optional(grouped_by) + pp.Suppress(')')
comparison = ((column | ARITHMETIC) + comparator + (quoted | column | ARITHMETIC)).set_parse_action(
# pyrefly: ignore[implicit-any-lambda]
lambda t: UnresolvedComparisonNode(t[0], t[1], t[2])
)
position_comparison = (position_call + comparator + position).set_parse_action(_position_comparison)

comparison = (column + comparator + (number | quoted | column)).set_parse_action(_comparison)
# pyrefly: ignore[implicit-any-lambda]
existence = name.copy().set_parse_action(lambda t: UnresolvedNameNode(t[0]))

atom = (
true_lit
| false_lit
| position_comparison
| comparison
| existence
| (pp.Suppress('(') + where_expr + pp.Suppress(')'))
)
atom = true_lit | false_lit | comparison | existence | (pp.Suppress('(') + where_expr + pp.Suppress(')'))

NOT = pp.CaselessKeyword('NOT').suppress()
# pyrefly: ignore[implicit-any-lambda]
Expand Down Expand Up @@ -199,6 +190,17 @@ def _named_rewrite(text: str, loc: int) -> str | None:
)


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, ArithmeticNode):
return children(node)
if isinstance(node, ConnectiveWhereNode):
return where_children(node)
return ()


@lru_cache(maxsize=4096)
def parse_where(text: str) -> WhereNode | UnresolvedWhereNode:
"""Parse a where string into an AST, its leaves still unresolved.
Expand All @@ -211,9 +213,10 @@ def parse_where(text: str) -> WhereNode | UnresolvedWhereNode:
SchemaError: If *text* is not a where string of the language. A
predictable mistake — ``&``/``|``/``~``/``!`` for a connective, a
lone ``=`` — is named with its rewrite beside the grammar's own
complaint.
complaint. A side nesting past what an expression may is refused
as an expression is.
"""
return cast(
'WhereNode | UnresolvedWhereNode',
parse_text(_WHERE_GRAMMAR, text, 'where string', _named_rewrite, where_children, _DEEP_REWRITE),
parse_text(_WHERE_GRAMMAR, text, 'where string', _named_rewrite, _nested, _DEEP_REWRITE),
)
Loading
Loading