368 lines
17 KiB
Python
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,
|
|
)
|