masque/masque/builder/tool_testing.py

368 lines
17 KiB
Python

"""Pytest-independent contract checks for custom routing Tools."""
from __future__ import annotations
from copy import deepcopy
from dataclasses import dataclass, field
from math import isfinite
from types import MappingProxyType
from typing import TYPE_CHECKING, Any
import numpy
from numpy import pi
from ..library import ILibrary, SINGLE_USE_PREFIX
from ..ports import Port
from ..utils import ptypes_compatible
from ._tolerances import angles_equal, array_close, scalar_close
from .error import ToolContractError
from .tools import (
BendOffer, PrimitiveKind, PrimitiveOffer, RenderStep, SOffer, StraightOffer, Tool, UOffer,
)
if TYPE_CHECKING:
from collections.abc import Mapping, Sequence
_RESERVED_OPTION_KEYS = frozenset(('kind', 'in_ptype', 'out_ptype', 'ccw'))
@dataclass(frozen=True, slots=True)
class ToolContractCase:
"""One primitive-discovery query to exercise against a custom Tool."""
kind: PrimitiveKind
in_ptype: str | None = None
out_ptype: str | None = None
ccw: bool | None = None
tool_options: Mapping[str, Any] = field(default_factory=lambda: MappingProxyType({}))
probe_parameters: tuple[float, ...] = ()
require_offers: bool = True
check_bbox: bool = False
label: str | None = None
def __post_init__(self) -> None:
if self.kind not in ('straight', 'bend', 's', 'u'):
raise ValueError(f'Unrecognized primitive kind {self.kind!r}')
if self.kind == 'bend':
if self.ccw is None:
raise ValueError('Bend ToolContractCase requires ccw')
elif self.ccw is not None:
raise ValueError('ccw is only valid for bend ToolContractCase')
try:
options = deepcopy(dict(self.tool_options))
except Exception as err:
raise ValueError('ToolContractCase.tool_options must be a deep-copyable mapping') from err
nonstring = [key for key in options if not isinstance(key, str)]
if nonstring:
raise ValueError(f'ToolContractCase.tool_options keys must be strings; got {nonstring!r}')
collisions = sorted(_RESERVED_OPTION_KEYS & options.keys())
if collisions:
raise ValueError(f'ToolContractCase.tool_options contains reserved keys: {", ".join(collisions)}')
try:
probes = tuple(float(value) for value in self.probe_parameters)
except (TypeError, ValueError, OverflowError) as err:
raise ValueError('ToolContractCase.probe_parameters must contain numeric scalars') from err
if not all(isfinite(value) for value in probes):
raise ValueError('ToolContractCase.probe_parameters must be finite')
object.__setattr__(self, 'ccw', None if self.ccw is None else bool(self.ccw))
object.__setattr__(self, 'tool_options', MappingProxyType(options))
object.__setattr__(self, 'probe_parameters', probes)
def _automatic_probes(offer: PrimitiveOffer) -> tuple[float, ...]:
"""Choose deterministic representative parameters inside one offer domain."""
lower, upper = (float(value) for value in offer.parameter_domain)
if lower == upper:
return (lower,)
if numpy.isfinite(lower) and numpy.isfinite(upper):
midpoint = lower / 2 + upper / 2
if midpoint == upper:
midpoint = float(numpy.nextafter(upper, lower))
return (lower, midpoint)
if numpy.isfinite(lower):
step = max(1.0, abs(lower) * 0.1)
return (lower, lower + step)
if numpy.isfinite(upper):
step = max(1.0, abs(upper) * 0.1)
return (upper - step, upper - 2 * step)
return (-1.0, 1.0)
def _offer_probes(offer: PrimitiveOffer, extras: Sequence[float]) -> tuple[float, ...]:
"""Combine automatic and applicable explicit probes without duplicates."""
probes: list[float] = []
for parameter in (*_automatic_probes(offer), *extras):
try:
selected = offer.canonicalize_parameter(parameter)
except Exception:
continue
if not any(scalar_close(selected, previous) for previous in probes):
probes.append(selected)
return tuple(probes)
def _offer_metadata(offer: PrimitiveOffer) -> tuple[Any, ...]:
"""Return discovery metadata that must remain stable across repeated queries."""
cost_policy: tuple[str, Any]
if callable(offer.cost):
cost_policy = ('callable', type(offer.cost).__qualname__)
else:
cost_policy = ('factor', float(offer.cost))
return (
type(offer),
offer.kind,
offer.in_ptype,
offer.out_ptype,
tuple(float(value) for value in offer.parameter_domain),
getattr(offer, 'ccw', None),
cost_policy,
)
def _evaluated_cost(offer: PrimitiveOffer, parameter: float, endpoint: Port) -> float:
"""Mirror the solver's one-endpoint base-cost path while honoring overrides."""
if type(offer).cost_at is PrimitiveOffer.cost_at:
return PrimitiveOffer._cost_for_endpoint(offer, parameter, endpoint)
return float(offer.cost_at(parameter))
def validate_tool_contract(tool: Tool, cases: Sequence[ToolContractCase]) -> None:
"""Validate Tool discovery, offer callbacks, and one-step rendering.
All independent violations are collected and raised as one
`ExceptionGroup` containing contextual `ToolContractError` instances.
"""
cases = tuple(cases)
if not cases:
raise ValueError('validate_tool_contract() requires at least one case')
errors: list[ToolContractError] = []
def violation(context: str, message: str, cause: Exception | None = None) -> None:
err = ToolContractError(f'{context}: {message}')
if cause is not None:
err.__cause__ = cause
errors.append(err)
def discover(case: ToolContractCase, context: str, repetition: str) -> tuple[PrimitiveOffer, ...] | None:
expected_offer_type = {
'straight': StraightOffer,
'bend': BendOffer,
's': SOffer,
'u': UOffer,
}[case.kind]
try:
kwargs = deepcopy(dict(case.tool_options))
if case.kind == 'bend':
kwargs['ccw'] = case.ccw
offers = tool.primitive_offers(
case.kind,
in_ptype=case.in_ptype,
out_ptype=case.out_ptype,
**kwargs,
)
except Exception as err:
violation(context, f'{repetition} discovery raised {type(err).__name__}: {err}', err)
return None
if not isinstance(offers, tuple):
violation(context, f'{repetition} discovery returned {type(offers).__name__}, expected tuple')
return None
valid = True
for offer_index, offer in enumerate(offers):
if not isinstance(offer, PrimitiveOffer):
violation(
context,
f'{repetition} discovery item {offer_index} is {type(offer).__name__}, expected PrimitiveOffer',
)
valid = False
elif offer.kind != case.kind:
violation(
context,
f'{repetition} discovery item {offer_index} has kind {offer.kind!r}, expected {case.kind!r}',
)
valid = False
elif not isinstance(offer, expected_offer_type):
violation(
context,
f'{repetition} discovery item {offer_index} is {type(offer).__name__}, '
f'expected {expected_offer_type.__name__}',
)
valid = False
return offers if valid else None
for case_index, case in enumerate(cases):
context = case.label or f'case {case_index} ({case.kind})'
first = discover(case, context, 'first')
second = discover(case, context, 'repeated')
if first is None or second is None:
continue
if case.require_offers and not first:
violation(context, 'discovery returned no offers')
if len(first) != len(second):
violation(context, f'discovery count changed from {len(first)} to {len(second)}')
matched_explicit = [False] * len(case.probe_parameters)
for offer_index, offer in enumerate(first):
offer_context = f'{context}, offer {offer_index}'
repeated = second[offer_index] if offer_index < len(second) else None
if repeated is not None and _offer_metadata(offer) != _offer_metadata(repeated):
violation(offer_context, 'discovery metadata changed between repeated queries')
probes = _offer_probes(offer, case.probe_parameters)
for explicit_index, parameter in enumerate(case.probe_parameters):
try:
offer.canonicalize_parameter(parameter)
except Exception:
continue
matched_explicit[explicit_index] = True
stable_ptype: str | None = None
stable_rotation: float | None = None
has_stable_endpoint = False
for parameter in probes:
probe_context = f'{offer_context}, parameter {parameter:g}'
try:
endpoint = offer.endpoint_at(parameter)
except Exception as err:
violation(probe_context, f'endpoint_at() raised {type(err).__name__}: {err}', err)
continue
if not isinstance(endpoint, Port):
violation(probe_context, f'endpoint_at() returned {type(endpoint).__name__}, expected Port')
continue
if not numpy.all(numpy.isfinite(endpoint.offset)):
violation(probe_context, 'endpoint offset must be finite')
if endpoint.rotation is None or not numpy.isfinite(endpoint.rotation):
violation(probe_context, 'endpoint rotation must be finite and specified')
if not ptypes_compatible(endpoint.ptype, offer.out_ptype):
violation(probe_context, 'endpoint ptype does not match declared out_ptype')
if offer.kind in ('straight', 'bend') and not scalar_close(endpoint.x, parameter):
violation(probe_context, 'straight/bend endpoint x must equal its parameter')
if offer.kind in ('s', 'u') and not scalar_close(endpoint.y, parameter):
violation(probe_context, 'S/U endpoint y must equal its parameter')
expected_rotation = {
'straight': pi,
'bend': -pi / 2 if isinstance(offer, BendOffer) and offer.ccw else pi / 2,
's': pi,
'u': 0.0,
}[offer.kind]
if endpoint.rotation is not None and not angles_equal(endpoint.rotation, expected_rotation):
violation(
probe_context,
f'endpoint rotation does not match {offer.kind!r} geometry',
)
if has_stable_endpoint:
if endpoint.ptype != stable_ptype:
violation(probe_context, 'endpoint ptype changes across the offer domain')
if (
endpoint.rotation is None
or stable_rotation is None
or not angles_equal(endpoint.rotation, stable_rotation)
):
violation(probe_context, 'endpoint rotation changes across the offer domain')
else:
stable_ptype = endpoint.ptype
stable_rotation = endpoint.rotation
has_stable_endpoint = True
try:
cost = float(_evaluated_cost(offer, parameter, endpoint))
if not numpy.isfinite(cost) or cost < 0:
violation(probe_context, f'cost must be finite and nonnegative, got {cost!r}')
except Exception as err:
violation(probe_context, f'cost_at() raised {type(err).__name__}: {err}', err)
cost = None
if repeated is not None:
try:
repeated_endpoint = repeated.endpoint_at(parameter)
repeated_cost = float(_evaluated_cost(repeated, parameter, repeated_endpoint))
if (
not isinstance(repeated_endpoint, Port)
or not array_close(repeated_endpoint.offset, endpoint.offset)
or repeated_endpoint.ptype != endpoint.ptype
or repeated_endpoint.rotation is None
or endpoint.rotation is None
or not angles_equal(repeated_endpoint.rotation, endpoint.rotation)
):
violation(probe_context, 'endpoint result changed after repeated discovery')
if cost is not None and not scalar_close(repeated_cost, cost):
violation(probe_context, 'cost result changed after repeated discovery')
except Exception as err:
violation(
probe_context,
f'repeated offer evaluation raised {type(err).__name__}: {err}',
err,
)
if case.check_bbox:
try:
offer.bbox_at(parameter)
except Exception as err:
violation(probe_context, f'bbox_at() raised {type(err).__name__}: {err}', err)
try:
data = offer.commit(parameter)
except Exception as err:
violation(probe_context, f'commit() raised {type(err).__name__}: {err}', err)
continue
try:
start = Port((0, 0), rotation=pi, ptype=offer.in_ptype or 'unk')
tree = tool.render((RenderStep(offer.kind, tool, start, endpoint.copy(), data),))
except Exception as err:
violation(probe_context, f'render() raised {type(err).__name__}: {err}', err)
continue
if not isinstance(tree, ILibrary):
violation(probe_context, f'render() returned {type(tree).__name__}, expected ILibrary')
continue
try:
top_name = tree.top()
pattern = tree.top_pattern()
except Exception as err:
violation(probe_context, f'rendered tree has no valid top cell: {err}', err)
continue
missing = sorted(
name
for name in tree.dangling_refs(top_name)
if isinstance(name, str) and name.startswith(SINGLE_USE_PREFIX)
)
if missing:
violation(probe_context, f'rendered tree has missing single-use refs: {missing}')
missing_ports = [name for name in ('A', 'B') if name not in pattern.ports]
if missing_ports:
violation(probe_context, f'rendered top cell is missing ports: {missing_ports}')
continue
input_port, output_port = pattern.ports['A'], pattern.ports['B']
if not ptypes_compatible(input_port.ptype, offer.in_ptype):
violation(probe_context, 'rendered input ptype does not match offer in_ptype')
try:
rendered_offset, rendered_rotation = input_port.measure_travel(output_port)
except Exception as err:
violation(probe_context, f'unable to measure rendered endpoint: {err}', err)
continue
if not array_close(rendered_offset, endpoint.offset):
violation(probe_context, 'rendered output offset does not match planned endpoint')
if (
rendered_rotation is None
or endpoint.rotation is None
or not angles_equal(rendered_rotation, endpoint.rotation)
):
violation(probe_context, 'rendered output rotation does not match planned endpoint')
if not ptypes_compatible(output_port.ptype, endpoint.ptype):
violation(probe_context, 'rendered output ptype does not match planned endpoint')
for parameter, matched in zip(case.probe_parameters, matched_explicit, strict=True):
if not matched:
violation(context, f'explicit probe {parameter:g} is outside every discovered offer domain')
if errors:
raise ExceptionGroup(
f'{type(tool).__name__} failed Tool contract validation with {len(errors)} violation(s)',
errors,
)