"""Typed intermediate representation for plane geometry diagrams.""" from __future__ import annotations from enum import Enum from typing import Annotated, Literal from pydantic import BaseModel, Field, model_validator class PointDef(BaseModel): name: str """Optional hint used only for schematic/initial placement.""" hint: tuple[float, float] | None = None class SegmentDef(BaseModel): a: str b: str draw: bool = True class LineDef(BaseModel): a: str b: str draw: bool = False """If true, draw a longer line through A and B rather than only the segment.""" class CircleDef(BaseModel): id: str center: str | None = None radius: float | None = None """Circle defined by three points on the circumference.""" through: list[str] | None = None draw: bool = True @model_validator(mode="after") def _check_definition(self) -> CircleDef: has_center = self.center is not None has_through = self.through is not None and len(self.through) >= 3 if not has_center and not has_through: raise ValueError(f"Circle {self.id!r} needs center or through points") if has_through and len(self.through or []) < 3: raise ValueError(f"Circle {self.id!r} through needs at least 3 points") return self class AngleMarkDef(BaseModel): """Mark angle at vertex with rays through a and b (order a-vertex-b).""" vertex: str a: str b: str right_angle: bool = False label: str | None = None class LabelDef(BaseModel): point: str text: str | None = None """TikZ anchor hint, e.g. above, below left.""" position: str = "above right" # --- Constraints --- class EqualLength(BaseModel): type: Literal["equal_length"] = "equal_length" a1: str a2: str b1: str b2: str class Length(BaseModel): type: Literal["length"] = "length" a: str b: str value: float class EqualAngle(BaseModel): type: Literal["equal_angle"] = "equal_angle" a1: str v1: str b1: str a2: str v2: str b2: str class AngleMeasure(BaseModel): type: Literal["angle_measure"] = "angle_measure" a: str vertex: str b: str degrees: float class Perpendicular(BaseModel): type: Literal["perpendicular"] = "perpendicular" a1: str a2: str b1: str b2: str class Parallel(BaseModel): type: Literal["parallel"] = "parallel" a1: str a2: str b1: str b2: str class OnLine(BaseModel): type: Literal["on_line"] = "on_line" point: str a: str b: str class OnCircle(BaseModel): type: Literal["on_circle"] = "on_circle" point: str circle: str class Midpoint(BaseModel): type: Literal["midpoint"] = "midpoint" point: str a: str b: str class Collinear(BaseModel): type: Literal["collinear"] = "collinear" points: list[str] @model_validator(mode="after") def _need_three(self) -> Collinear: if len(self.points) < 3: raise ValueError("collinear needs at least 3 points") return self class Intersection(BaseModel): """Point is intersection of lines AB and CD.""" type: Literal["intersection"] = "intersection" point: str a: str b: str c: str d: str Constraint = Annotated[ EqualLength | Length | EqualAngle | AngleMeasure | Perpendicular | Parallel | OnLine | OnCircle | Midpoint | Collinear | Intersection, Field(discriminator="type"), ] class GeometryIR(BaseModel): """Formal scene: objects, constraints, and draw directives.""" statement: str = "" points: list[PointDef] segments: list[SegmentDef] = Field(default_factory=list) lines: list[LineDef] = Field(default_factory=list) circles: list[CircleDef] = Field(default_factory=list) angle_marks: list[AngleMarkDef] = Field(default_factory=list) labels: list[LabelDef] = Field(default_factory=list) constraints: list[Constraint] = Field(default_factory=list) def point_names(self) -> list[str]: return [p.name for p in self.points] def ensure_labels(self) -> None: existing = {lb.point for lb in self.labels} for p in self.points: if p.name not in existing: self.labels.append(LabelDef(point=p.name, text=p.name)) class SolveMode(str, Enum): exact = "exact" schematic = "schematic" failed = "failed" class SolvedScene(BaseModel): ir: GeometryIR coordinates: dict[str, tuple[float, float]] mode: SolveMode max_residual: float residuals: list[float] = Field(default_factory=list) message: str = ""