Files
reducers/backend/src/relation_validators/base.py
T

187 lines
4.9 KiB
Python

from __future__ import annotations
import math
from dataclasses import dataclass
from typing import Any, Protocol
from ..models import (
AssemblySnapshot,
ComponentManifest,
ConcreteNode,
ConcreteRelation,
ConcreteTopology,
ReducerInstance,
ValidationCheck,
)
@dataclass(frozen=True)
class RelationValidationContext:
instance: ReducerInstance
topology: ConcreteTopology
snapshot: AssemblySnapshot
def node(self, component_id: str) -> ConcreteNode | None:
return next(
(node for node in self.topology.nodes if node.component_id == component_id),
None,
)
def component(self, component_id: str) -> ComponentManifest | None:
return self.snapshot.manifest.components.get(component_id)
def placement_origin(self, component_id: str) -> list[float] | None:
origin = self.snapshot.component_placements.get(component_id, {}).get("origin")
if origin is None:
component = self.component(component_id)
if component is None:
return None
origin = component.center_xyz_mm
return [float(value) for value in origin]
class RelationValidator(Protocol):
relation_type: str
def validate_static(
self,
context: RelationValidationContext,
relation: ConcreteRelation,
checks: list[ValidationCheck],
) -> None:
...
def validate_assembly(
self,
context: RelationValidationContext,
relation: ConcreteRelation,
checks: list[ValidationCheck],
) -> None:
...
def validate_geometry(
self,
context: RelationValidationContext,
relation: ConcreteRelation,
checks: list[ValidationCheck],
) -> None:
...
class BaseRelationValidator:
relation_type: str
def validate_static(
self,
context: RelationValidationContext,
relation: ConcreteRelation,
checks: list[ValidationCheck],
) -> None:
return None
def validate_assembly(
self,
context: RelationValidationContext,
relation: ConcreteRelation,
checks: list[ValidationCheck],
) -> None:
return None
def validate_geometry(
self,
context: RelationValidationContext,
relation: ConcreteRelation,
checks: list[ValidationCheck],
) -> None:
return None
def add_check(
checks: list[ValidationCheck],
*,
section: str,
code: str,
passed: bool,
message: str,
actual: Any | None = None,
expected: Any | None = None,
) -> None:
checks.append(
ValidationCheck(
section=section,
code=code,
passed=bool(passed),
message=message,
actual=actual,
expected=expected,
)
)
def relation_code(relation: ConcreteRelation, suffix: str) -> str:
return f"{relation.relation_type}.{relation.relation_id}.{suffix}"
def xy_distance(origin_a: list[float], origin_b: list[float]) -> float:
return math.hypot(origin_a[0] - origin_b[0], origin_a[1] - origin_b[1])
def xy_radius(origin: list[float]) -> float:
return math.hypot(origin[0], origin[1])
def angle_deg(origin: list[float]) -> float:
return (math.degrees(math.atan2(origin[1], origin[0])) + 360.0) % 360.0
def find_manifest_constraint(
snapshot: AssemblySnapshot,
*,
relation_type: str,
component_a: str,
component_b: str,
allow_reversed: bool = True,
) -> str | None:
for constraint in snapshot.manifest.constraints:
if constraint.relation_type != relation_type:
continue
exact = constraint.component_a == component_a and constraint.component_b == component_b
reversed_pair = (
allow_reversed
and constraint.component_a == component_b
and constraint.component_b == component_a
)
if exact or reversed_pair:
return constraint.constraint_id
return None
def constraint_id_in_model(
snapshot: AssemblySnapshot,
*,
constraint_id: str | None,
relation_type: str,
) -> bool:
if not constraint_id:
return False
lists = {
"external_mesh": snapshot.gear_constraints,
"internal_mesh": snapshot.internal_mesh_constraints,
"revolute_joint": snapshot.revolute_constraints,
"fixed": snapshot.fixed_constraints,
"rigid": snapshot.fixed_constraints,
}
return constraint_id in lists.get(relation_type, [])
def same_axis(axis_a: tuple[float, float, float], axis_b: tuple[float, float, float]) -> bool:
ax = tuple(float(value) for value in axis_a)
bx = tuple(float(value) for value in axis_b)
dot = sum(a * b for a, b in zip(ax, bx))
mag_a = math.sqrt(sum(a * a for a in ax))
mag_b = math.sqrt(sum(b * b for b in bx))
if mag_a <= 1e-12 or mag_b <= 1e-12:
return False
return abs(abs(dot / (mag_a * mag_b)) - 1.0) <= 1e-9