[planner] make solver carry state across increasing bend counts

This commit is contained in:
Jan Petykiewicz 2026-07-09 11:09:52 -07:00
commit f864ebbeab
2 changed files with 186 additions and 101 deletions

View file

@ -293,16 +293,24 @@ class Solver:
tuple[PrimitiveKind, str | None, str | None, tuple[tuple[str, Any], ...]],
tuple[PrimitiveOffer, ...],
] = {}
self.seen_candidate_keys: set[tuple[Any, ...]] = set()
self.order = 0
def solve(self) -> Candidate:
"""
Enumerate, finalize, deduplicate, and rank legal candidates.
@staticmethod
def route_bend_count(steps: Sequence[SelectedPrimitive]) -> int:
"""Return the route bend budget consumed by non-adapter primitives."""
count = 0
for step in steps:
if step.role == 'adapter':
continue
if step.route_kind == 'bend':
count += 1
elif step.route_kind in ('s', 'u'):
count += 2
return count
Non-fatal candidate errors are accumulated so the failure message can
preserve useful Tool feedback. Fatal offer-contract errors stop the
solve immediately.
"""
def candidate_key(self, candidate: Candidate) -> tuple[Any, ...]:
"""Return a deterministic key for duplicate solved candidates."""
def endpoint_key(port: Port) -> tuple[float, float, float | None, str | None]:
return (
round(float(port.x), 9),
@ -311,36 +319,51 @@ class Solver:
port.ptype,
)
def candidate_key(candidate: Candidate) -> tuple[Any, ...]:
def offer_key(offer: PrimitiveOffer) -> tuple[Any, ...]:
return (
type(offer).__qualname__,
offer.in_ptype,
offer.out_ptype,
round(float(offer.priority_bias), 9),
tuple(round(float(value), 9) for value in offer.parameter_domain),
getattr(offer, 'ccw', None),
id(offer.endpoint_planner),
id(offer.commit_planner),
)
def offer_key(offer: PrimitiveOffer) -> tuple[Any, ...]:
return (
endpoint_key(candidate.end_port),
tuple((
offer_key(step.offer),
step.role,
step.route_kind,
round(float(step.parameter), 9),
endpoint_key(step.out_port),
) for step in candidate.steps),
type(offer).__qualname__,
offer.in_ptype,
offer.out_ptype,
round(float(offer.priority_bias), 9),
tuple(round(float(value), 9) for value in offer.parameter_domain),
getattr(offer, 'ccw', None),
id(offer.endpoint_planner),
id(offer.commit_planner),
)
return (
endpoint_key(candidate.end_port),
tuple((
offer_key(step.offer),
step.role,
step.route_kind,
round(float(step.parameter), 9),
endpoint_key(step.out_port),
) for step in candidate.steps),
)
def solve(
self,
*,
min_bends: int = 0,
max_bends: int | None = None,
) -> Candidate:
"""
Enumerate, finalize, deduplicate, and rank legal candidates.
Non-fatal candidate errors are accumulated so the failure message can
preserve useful Tool feedback. Fatal offer-contract errors stop the
solve immediately.
"""
if max_bends is None:
max_bends = self.request.bend_budget
candidates: list[Candidate] = []
errors: list[Exception] = []
seen: set[tuple[Any, ...]] = set()
for steps in self.enumerate_grammar():
for steps in self.enumerate_grammar(max_bends):
if not steps:
continue
if self.route_bend_count(steps) < min_bends:
continue
if any(first.role == 'adapter' and second.role == 'adapter' for first, second in zip(steps, steps[1:], strict=False)):
continue
try:
@ -349,10 +372,10 @@ class Solver:
raise_if_fatal(err)
errors.append(err)
continue
key = candidate_key(candidate)
if key in seen:
key = self.candidate_key(candidate)
if key in self.seen_candidate_keys:
continue
seen.add(key)
self.seen_candidate_keys.add(key)
candidates.append(candidate)
if not candidates:
@ -676,13 +699,13 @@ class Solver:
options.extend((turn, 2) for turn in self.su_primitive_options(steps, 'u', jog))
return tuple(options)
def enumerate_grammar(self) -> Iterable[tuple[SelectedPrimitive, ...]]:
def enumerate_grammar(self, max_bends: int) -> Iterable[tuple[SelectedPrimitive, ...]]:
"""Yield raw primitive sequences allowed by the bounded route grammar."""
residual_jog = 0.0 if self.request.jog is None else float(self.request.jog)
base: tuple[SelectedPrimitive, ...] = ()
for prefix_adapter in self.adapter_options(base, residual_jog=residual_jog):
prefix = (*base, *prefix_adapter)
yield from self.enumerate_segments(prefix, self.request.bend_budget, residual_jog=residual_jog)
yield from self.enumerate_segments(prefix, max_bends, residual_jog=residual_jog)
def enumerate_segments(
self,
@ -909,11 +932,63 @@ class RoutingPlanner:
TRACE_INTO_MAX_BENDS: int = 4
def trace_into_bend_budgets(self, family: PrimitiveKind) -> tuple[int, ...]:
"""Return staged trace_into bend budgets for the route rotation parity."""
def trace_into_bend_bands(self, family: PrimitiveKind) -> tuple[tuple[int, int], ...]:
"""Return non-overlapping trace_into bend-budget bands for staged solving."""
max_bends = self.TRACE_INTO_MAX_BENDS
if family == 'bend':
return tuple(budget for budget in (1, 3) if budget <= self.TRACE_INTO_MAX_BENDS)
return tuple(budget for budget in (2, 4) if budget <= self.TRACE_INTO_MAX_BENDS)
return tuple(band for band in ((1, 1), (3, 3)) if band[1] <= max_bends)
bands: list[tuple[int, int]] = []
if max_bends >= 0:
bands.append((0, 2 if max_bends >= 2 else 0))
if max_bends >= 4:
bands.append((4, 4))
return tuple(bands)
def route_request(
self,
family: PrimitiveKind,
context: RoutePortContext,
*,
length: float | None = None,
jog: float | None = None,
ccw: SupportsBool | None = None,
constrain_jog: bool = False,
max_bends: int | None = None,
**kwargs: Any,
) -> RouteRequest:
"""Build normalized solver input for one route leg."""
return RouteRequest(
family=family,
tool=context.tool,
in_ptype=context.port.ptype,
route_kwargs=kwargs,
length=length,
jog=jog,
ccw=ccw,
out_ptype=kwargs.get('out_ptype'),
constrain_jog=constrain_jog,
max_bends=max_bends,
)
def solver_for_request(self, request: RouteRequest) -> Solver:
"""Construct the solver for a route request."""
return Solver(request)
def route_leg_from_candidate(
self,
context: RoutePortContext,
candidate: Candidate,
*,
plug_into: str | None,
) -> RouteLeg:
"""Attach a solved candidate to its source Pather context."""
return RouteLeg(
portspec=context.portspec,
start_port=context.port.copy(),
tool=context.tool,
candidate=candidate,
plug_into=plug_into,
)
def plan_leg(
self,
@ -929,31 +1004,23 @@ class RoutingPlanner:
**kwargs: Any,
) -> RouteLeg:
"""Solve one route leg and attach it to its source Pather context."""
request = RouteRequest(
request = self.route_request(
family=family,
tool=context.tool,
in_ptype=context.port.ptype,
route_kwargs=kwargs,
context=context,
length=length,
jog=jog,
ccw=ccw,
out_ptype=kwargs.get('out_ptype'),
constrain_jog=constrain_jog,
max_bends=max_bends,
**kwargs,
)
try:
candidate = Solver(request).solve()
candidate = self.solver_for_request(request).solve()
except BuildError as err:
if family == 'u' and length is None and not getattr(err, 'fatal', False):
raise BuildError('No legal primitive offer for omitted-length U-turn') from err
raise
return RouteLeg(
portspec=context.portspec,
start_port=context.port.copy(),
tool=context.tool,
candidate=candidate,
plug_into=plug_into,
)
return self.route_leg_from_candidate(context, candidate, plug_into=plug_into)
def prepared_route_action_from_leg(
self,
@ -1204,30 +1271,36 @@ class RoutingPlanner:
desired.rotation = port_dst.rotation - pi
desired.ptype = out_ptype
family, length, jog, ccw = self.trace_into_spec(context_src.port, desired)
leg = None
request = self.route_request(
family,
context_src,
length=length,
jog=jog,
ccw=ccw,
constrain_jog=family == 'bend',
max_bends=self.TRACE_INTO_MAX_BENDS,
**(dict(kwargs) | {'out_ptype': out_ptype}),
)
solver = self.solver_for_request(request)
candidate = None
last_error: Exception | None = None
for max_bends in self.trace_into_bend_budgets(family):
for min_bends, max_bends in self.trace_into_bend_bands(family):
try:
leg = self.plan_leg(
family,
context_src,
length=length,
jog=jog,
ccw=ccw,
plug_into=portspec_dst if plug_destination else None,
constrain_jog=family == 'bend',
max_bends=max_bends,
**(dict(kwargs) | {'out_ptype': out_ptype}),
)
candidate = solver.solve(min_bends=min_bends, max_bends=max_bends)
break
except (BuildError, NotImplementedError) as err:
if route_error_is_fatal(err):
raise
last_error = err
if leg is None:
if candidate is None:
if last_error is not None:
raise last_error
raise BuildError('No legal primitive offer for trace_into route')
leg = self.route_leg_from_candidate(
context_src,
candidate,
plug_into=portspec_dst if plug_destination else None,
)
renames = ((thru, context_src.portspec),) if thru is not None else ()
return self.prepared_result_from_legs(
(leg,),