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