Skip to main content

quoracle/
geometry.rs

1//! Geometric types for piecewise linear functions.
2//!
3//! Provides [`Point`], [`Segment`], and [`max_of_segments`] for
4//! computing upper envelopes of line segments on \[0, 1\].
5
6use crate::error::{Error, Result};
7
8/// A point in 2D space.
9#[derive(Debug, Clone, Copy, PartialEq)]
10pub struct Point {
11    /// X coordinate.
12    pub x: f64,
13    /// Y coordinate.
14    pub y: f64,
15}
16
17impl Point {
18    /// Create a new point.
19    #[must_use]
20    pub const fn new(x: f64, y: f64) -> Self {
21        Self { x, y }
22    }
23}
24
25/// A line segment between two points where `l.x < r.x`.
26#[derive(Debug, Clone, Copy)]
27pub struct Segment {
28    /// Left endpoint (smaller x).
29    pub l: Point,
30    /// Right endpoint (larger x).
31    pub r: Point,
32}
33
34impl PartialEq for Segment {
35    fn eq(&self, other: &Self) -> bool {
36        self.l == other.l && self.r == other.r
37    }
38}
39
40impl Segment {
41    /// Create a new segment from left point `l` to right point `r`.
42    ///
43    /// # Errors
44    ///
45    /// Returns an error if `l == r` or `l.x >= r.x`.
46    pub fn new(l: Point, r: Point) -> Result<Self> {
47        if l == r {
48            return Err(Error::InvalidExpression(
49                "segment endpoints must differ".into(),
50            ));
51        }
52        if l.x >= r.x {
53            return Err(Error::InvalidExpression(
54                "left endpoint x must be less than right endpoint x".into(),
55            ));
56        }
57        Ok(Self { l, r })
58    }
59
60    /// Evaluate the linear function at `x`.
61    ///
62    /// # Errors
63    ///
64    /// Returns an error if `x` is outside the segment's x-range.
65    pub fn eval(&self, x: f64) -> Result<f64> {
66        if x < self.l.x || x > self.r.x {
67            return Err(Error::InvalidExpression(format!(
68                "x={x} is outside segment range [{}, {}]",
69                self.l.x, self.r.x
70            )));
71        }
72        Ok(self.slope() * (x - self.l.x) + self.l.y)
73    }
74
75    /// Whether two segments are approximately equal (within relative
76    /// tolerance 1e-5 on both y-coordinates).
77    #[must_use]
78    pub fn approximately_equal(&self, other: &Self) -> bool {
79        approx_eq(self.l.y, other.l.y) && approx_eq(self.r.y, other.r.y)
80    }
81
82    /// Whether two segments share the same x-range.
83    #[must_use]
84    #[expect(clippy::float_cmp)]
85    pub fn compatible(&self, other: &Self) -> bool {
86        self.l.x == other.l.x && self.r.x == other.r.x
87    }
88
89    /// The slope `(r.y - l.y) / (r.x - l.x)`.
90    #[must_use]
91    pub fn slope(&self) -> f64 {
92        (self.r.y - self.l.y) / (self.r.x - self.l.x)
93    }
94
95    /// Whether `self` is strictly above `other` (both endpoints at
96    /// least as high, and not equal).
97    ///
98    /// # Errors
99    ///
100    /// Returns an error if the segments are not compatible.
101    pub fn above(&self, other: &Self) -> Result<bool> {
102        self.assert_compatible(other)?;
103        Ok(self != other && self.l.y >= other.l.y && self.r.y >= other.r.y)
104    }
105
106    /// Whether `self` is above or equal to `other`.
107    ///
108    /// # Errors
109    ///
110    /// Returns an error if the segments are not compatible.
111    pub fn above_eq(&self, other: &Self) -> Result<bool> {
112        self.assert_compatible(other)?;
113        Ok(self == other || self.above_unchecked(other))
114    }
115
116    /// Whether two compatible segments intersect.
117    ///
118    /// # Errors
119    ///
120    /// Returns an error if the segments are not compatible.
121    pub fn intersects(&self, other: &Self) -> Result<bool> {
122        self.assert_compatible(other)?;
123        Ok(self.intersects_unchecked(other))
124    }
125
126    /// Compute the intersection point of two compatible segments, or
127    /// `None` if they are equal or do not intersect.
128    ///
129    /// The x-coordinate formula assumes both segments span \[0, 1\].
130    ///
131    /// # Errors
132    ///
133    /// Returns an error if the segments are not compatible.
134    pub fn intersection(&self, other: &Self) -> Result<Option<Point>> {
135        self.assert_compatible(other)?;
136        if self == other || !self.intersects_unchecked(other) {
137            return Ok(None);
138        }
139        let denom = self.r.y - other.r.y + other.l.y - self.l.y;
140        let x = (other.l.y - self.l.y) / denom;
141        // x is guaranteed within [l.x, r.x] for intersecting
142        // compatible segments, so we compute y inline.
143        let y = self.slope() * (x - self.l.x) + self.l.y;
144        Ok(Some(Point::new(x, y)))
145    }
146
147    fn assert_compatible(&self, other: &Self) -> Result<()> {
148        if self.compatible(other) {
149            Ok(())
150        } else {
151            Err(Error::InvalidExpression(
152                "segments are not compatible (different x-ranges)".into(),
153            ))
154        }
155    }
156
157    fn above_unchecked(&self, other: &Self) -> bool {
158        self != other && self.l.y >= other.l.y && self.r.y >= other.r.y
159    }
160
161    #[expect(clippy::float_cmp)]
162    fn intersects_unchecked(&self, other: &Self) -> bool {
163        if self == other {
164            return true;
165        }
166        if self.l.y == other.l.y || self.r.y == other.r.y {
167            return true;
168        }
169        if self.above_unchecked(other) || other.above_unchecked(self) {
170            return false;
171        }
172        true
173    }
174}
175
176/// Compute the upper envelope of a set of compatible segments.
177///
178/// Returns a list of `(x, y)` pairs tracing the maximum of all
179/// segments. All segments must share the same x-range (typically
180/// \[0, 1\]).
181///
182/// # Errors
183///
184/// Returns an error if the slice is empty or segments have
185/// different x-ranges.
186#[expect(clippy::float_cmp)]
187pub fn max_of_segments(segments: &[Segment]) -> Result<Vec<(f64, f64)>> {
188    if segments.is_empty() {
189        return Err(Error::InvalidExpression(
190            "max_of_segments requires at least one segment".into(),
191        ));
192    }
193
194    let l_x = segments[0].l.x;
195    let r_x = segments[0].r.x;
196    for s in &segments[1..] {
197        if s.l.x != l_x || s.r.x != r_x {
198            return Err(Error::InvalidExpression(
199                "all segments must have the same x-range".into(),
200            ));
201        }
202    }
203
204    // Collect x-coordinates of all intersection points plus
205    // endpoints.
206    let mut xs: Vec<f64> = vec![0.0, 1.0];
207    for (i, s1) in segments.iter().enumerate() {
208        for s2 in &segments[i + 1..] {
209            if let Some(p) = s1.intersection(s2)? {
210                xs.push(p.x);
211            }
212        }
213    }
214    xs.sort_by(f64::total_cmp);
215
216    let mut result = Vec::with_capacity(xs.len());
217    for x in xs {
218        let mut max_y = f64::NEG_INFINITY;
219        for s in segments {
220            let y = s.eval(x)?;
221            if y > max_y {
222                max_y = y;
223            }
224        }
225        result.push((x, max_y));
226    }
227    Ok(result)
228}
229
230/// Relative tolerance 1e-5 (like Python's `math.isclose(rel_tol=1e-5)`).
231fn approx_eq(a: f64, b: f64) -> bool {
232    (a - b).abs() <= 1e-5 * a.abs().max(b.abs())
233}
234
235#[cfg(test)]
236#[expect(clippy::expect_used)]
237mod tests {
238    use super::*;
239
240    fn pt(x: f64, y: f64) -> Point {
241        Point::new(x, y)
242    }
243
244    fn seg(lx: f64, ly: f64, rx: f64, ry: f64) -> Segment {
245        Segment::new(pt(lx, ly), pt(rx, ry)).expect("valid segment")
246    }
247
248    #[test]
249    fn test_eq() {
250        let l = pt(0.0, 1.0);
251        let r = pt(1.0, 1.0);
252        let m = pt(0.5, 0.5);
253        assert_eq!(
254            Segment::new(l, r).expect("ok"),
255            Segment::new(l, r).expect("ok")
256        );
257        assert_ne!(
258            Segment::new(l, r).expect("ok"),
259            Segment::new(l, m).expect("ok")
260        );
261    }
262
263    #[test]
264    fn test_compatible() {
265        let s1 = seg(0.0, 1.0, 1.0, 2.0);
266        let s2 = seg(0.0, 2.0, 1.0, 1.0);
267        let s3 = seg(0.5, 2.0, 1.0, 1.0);
268        assert!(s1.compatible(&s2));
269        assert!(s2.compatible(&s1));
270        assert!(!s1.compatible(&s3));
271        assert!(!s3.compatible(&s1));
272        assert!(!s2.compatible(&s3));
273        assert!(!s3.compatible(&s2));
274    }
275
276    #[test]
277    fn test_eval() {
278        let segment = seg(0.0, 0.0, 1.0, 1.0);
279        for &x in &[0.0, 0.25, 0.5, 0.75, 1.0] {
280            assert_eq!(segment.eval(x).expect("ok"), x);
281        }
282
283        let segment = seg(0.0, 0.0, 1.0, 2.0);
284        for &x in &[0.0, 0.25, 0.5, 0.75, 1.0] {
285            assert_eq!(segment.eval(x).expect("ok"), 2.0 * x);
286        }
287
288        let segment = seg(1.0, 2.0, 3.0, 6.0);
289        for &x in &[1.0, 1.25, 1.5, 1.75, 2.0, 2.25, 2.5, 2.75, 3.0] {
290            assert_eq!(segment.eval(x).expect("ok"), 2.0 * x);
291        }
292
293        let segment = seg(0.0, 1.0, 1.0, 0.0);
294        for &x in &[0.0, 0.25, 0.5, 0.75, 1.0] {
295            assert_eq!(segment.eval(x).expect("ok"), 1.0 - x);
296        }
297    }
298
299    #[test]
300    fn test_slope() {
301        assert_eq!(seg(0.0, 0.0, 1.0, 1.0).slope(), 1.0);
302        assert_eq!(seg(0.0, 1.0, 1.0, 2.0).slope(), 1.0);
303        assert_eq!(seg(1.0, 1.0, 2.0, 2.0).slope(), 1.0);
304        assert_eq!(seg(1.0, 1.0, 2.0, 3.0).slope(), 2.0);
305        assert_eq!(seg(1.0, 1.0, 2.0, 0.0).slope(), -1.0);
306    }
307
308    #[test]
309    fn test_above() {
310        let s1 = seg(0.0, 0.0, 1.0, 0.5);
311        let s2 = seg(0.0, 0.5, 1.0, 2.0);
312        let s3 = seg(0.0, 1.5, 1.0, 0.5);
313
314        assert!(!s1.above(&s1).expect("ok"));
315        assert!(!s1.above(&s2).expect("ok"));
316        assert!(!s1.above(&s3).expect("ok"));
317
318        assert!(s2.above(&s1).expect("ok"));
319        assert!(!s2.above(&s2).expect("ok"));
320        assert!(!s2.above(&s3).expect("ok"));
321
322        assert!(s3.above(&s1).expect("ok"));
323        assert!(!s3.above(&s2).expect("ok"));
324        assert!(!s3.above(&s3).expect("ok"));
325    }
326
327    #[test]
328    fn test_above_eq() {
329        let s1 = seg(0.0, 0.0, 1.0, 0.5);
330        let s2 = seg(0.0, 0.5, 1.0, 2.0);
331        let s3 = seg(0.0, 1.5, 1.0, 0.5);
332
333        assert!(s1.above_eq(&s1).expect("ok"));
334        assert!(!s1.above_eq(&s2).expect("ok"));
335        assert!(!s1.above_eq(&s3).expect("ok"));
336
337        assert!(s2.above_eq(&s1).expect("ok"));
338        assert!(s2.above_eq(&s2).expect("ok"));
339        assert!(!s2.above_eq(&s3).expect("ok"));
340
341        assert!(s3.above_eq(&s1).expect("ok"));
342        assert!(!s3.above_eq(&s2).expect("ok"));
343        assert!(s3.above_eq(&s3).expect("ok"));
344    }
345
346    #[test]
347    fn test_intersects() {
348        let s1 = seg(0.0, 0.0, 1.0, 0.5);
349        let s2 = seg(0.0, 0.5, 1.0, 2.0);
350        let s3 = seg(0.0, 1.5, 1.0, 0.5);
351
352        assert!(s1.intersects(&s1).expect("ok"));
353        assert!(!s1.intersects(&s2).expect("ok"));
354        assert!(s1.intersects(&s3).expect("ok"));
355
356        assert!(!s2.intersects(&s1).expect("ok"));
357        assert!(s2.intersects(&s2).expect("ok"));
358        assert!(s2.intersects(&s3).expect("ok"));
359
360        assert!(s3.intersects(&s1).expect("ok"));
361        assert!(s3.intersects(&s2).expect("ok"));
362        assert!(s3.intersects(&s3).expect("ok"));
363    }
364
365    #[test]
366    fn test_intersection() {
367        let s1 = seg(0.0, 0.0, 1.0, 1.0);
368        let s2 = seg(0.0, 1.0, 1.0, 0.0);
369        let s3 = seg(0.0, 1.0, 1.0, 1.0);
370        let s4 = seg(0.0, 0.25, 1.0, 0.25);
371
372        assert_eq!(s1.intersection(&s1).expect("ok"), None);
373        assert_eq!(s1.intersection(&s2).expect("ok"), Some(pt(0.5, 0.5)));
374        assert_eq!(s1.intersection(&s3).expect("ok"), Some(pt(1.0, 1.0)));
375        assert_eq!(s1.intersection(&s4).expect("ok"), Some(pt(0.25, 0.25)));
376
377        assert_eq!(s2.intersection(&s1).expect("ok"), Some(pt(0.5, 0.5)));
378        assert_eq!(s2.intersection(&s2).expect("ok"), None);
379        assert_eq!(s2.intersection(&s3).expect("ok"), Some(pt(0.0, 1.0)));
380        assert_eq!(s2.intersection(&s4).expect("ok"), Some(pt(0.75, 0.25)));
381
382        assert_eq!(s3.intersection(&s1).expect("ok"), Some(pt(1.0, 1.0)));
383        assert_eq!(s3.intersection(&s2).expect("ok"), Some(pt(0.0, 1.0)));
384        assert_eq!(s3.intersection(&s3).expect("ok"), None);
385        assert_eq!(s3.intersection(&s4).expect("ok"), None);
386
387        assert_eq!(s4.intersection(&s1).expect("ok"), Some(pt(0.25, 0.25)));
388        assert_eq!(s4.intersection(&s2).expect("ok"), Some(pt(0.75, 0.25)));
389        assert_eq!(s4.intersection(&s3).expect("ok"), None);
390        assert_eq!(s4.intersection(&s4).expect("ok"), None);
391    }
392
393    #[test]
394    fn test_max_one_segment() {
395        let s1 = seg(0.0, 0.0, 1.0, 1.0);
396        let s2 = seg(0.0, 1.0, 1.0, 0.0);
397        let s3 = seg(0.0, 1.0, 1.0, 1.0);
398        let s4 = seg(0.0, 0.25, 1.0, 0.25);
399        let s5 = seg(0.0, 0.75, 1.0, 0.75);
400
401        for s in &[s1, s2, s3, s4, s5] {
402            let result = max_of_segments(&[*s]).expect("ok");
403            assert_eq!(result, vec![(s.l.x, s.l.y), (s.r.x, s.r.y)]);
404        }
405    }
406
407    fn is_subset(xs: &[(f64, f64)], ys: &[(f64, f64)]) -> bool {
408        xs.iter().all(|x| ys.contains(x))
409    }
410
411    type SegmentCase = (Vec<Segment>, Vec<(f64, f64)>);
412
413    #[test]
414    fn test_max_two_segments() {
415        let s1 = seg(0.0, 0.0, 1.0, 1.0);
416        let s2 = seg(0.0, 1.0, 1.0, 0.0);
417        let s3 = seg(0.0, 1.0, 1.0, 1.0);
418        let s4 = seg(0.0, 0.25, 1.0, 0.25);
419        let s5 = seg(0.0, 0.75, 1.0, 0.75);
420
421        let cases: Vec<SegmentCase> = vec![
422            (vec![s1, s1], vec![(0.0, 0.0), (1.0, 1.0)]),
423            (vec![s1, s2], vec![(0.0, 1.0), (0.5, 0.5), (1.0, 1.0)]),
424            (vec![s1, s3], vec![(0.0, 1.0), (1.0, 1.0)]),
425            (vec![s1, s4], vec![(0.0, 0.25), (0.25, 0.25), (1.0, 1.0)]),
426            (vec![s1, s5], vec![(0.0, 0.75), (0.75, 0.75), (1.0, 1.0)]),
427            (vec![s2, s2], vec![(0.0, 1.0), (1.0, 0.0)]),
428            (vec![s2, s3], vec![(0.0, 1.0), (1.0, 1.0)]),
429            (vec![s2, s4], vec![(0.0, 1.0), (0.75, 0.25), (1.0, 0.25)]),
430            (vec![s2, s5], vec![(0.0, 1.0), (0.25, 0.75), (1.0, 0.75)]),
431            (vec![s3, s3], vec![(0.0, 1.0), (1.0, 1.0)]),
432            (vec![s3, s4], vec![(0.0, 1.0), (1.0, 1.0)]),
433            (vec![s3, s5], vec![(0.0, 1.0), (1.0, 1.0)]),
434            (vec![s4, s4], vec![(0.0, 0.25), (1.0, 0.25)]),
435            (vec![s4, s5], vec![(0.0, 0.75), (1.0, 0.75)]),
436            (vec![s5, s5], vec![(0.0, 0.75), (1.0, 0.75)]),
437        ];
438
439        for (segments, path) in &cases {
440            let result = max_of_segments(segments).expect("ok");
441            assert!(
442                is_subset(path, &result),
443                "forward: {path:?} not subset of {result:?}"
444            );
445            let reversed: Vec<Segment> =
446                segments.iter().rev().copied().collect();
447            let result_rev = max_of_segments(&reversed).expect("ok");
448            assert!(
449                is_subset(path, &result_rev),
450                "reverse: {path:?} not subset of {result_rev:?}"
451            );
452        }
453    }
454
455    #[test]
456    fn test_max_three_segments() {
457        let s1 = seg(0.0, 0.0, 1.0, 1.0);
458        let s2 = seg(0.0, 1.0, 1.0, 0.0);
459        let s4 = seg(0.0, 0.25, 1.0, 0.25);
460        let s5 = seg(0.0, 0.75, 1.0, 0.75);
461
462        let cases: Vec<SegmentCase> = vec![
463            (vec![s1, s2, s4], vec![(0.0, 1.0), (0.5, 0.5), (1.0, 1.0)]),
464            (
465                vec![s1, s2, s5],
466                vec![(0.0, 1.0), (0.25, 0.75), (0.75, 0.75), (1.0, 1.0)],
467            ),
468        ];
469
470        for (segments, path) in &cases {
471            let result = max_of_segments(segments).expect("ok");
472            assert!(
473                is_subset(path, &result),
474                "forward: {path:?} not subset of {result:?}"
475            );
476            let reversed: Vec<Segment> =
477                segments.iter().rev().copied().collect();
478            let result_rev = max_of_segments(&reversed).expect("ok");
479            assert!(
480                is_subset(path, &result_rev),
481                "reverse: {path:?} not subset of {result_rev:?}"
482            );
483        }
484    }
485
486    #[test]
487    fn test_new_segment_invalid() {
488        let p = pt(1.0, 1.0);
489        assert!(Segment::new(p, p).is_err());
490        assert!(Segment::new(pt(1.0, 0.0), pt(0.0, 1.0)).is_err());
491    }
492
493    #[test]
494    fn test_eval_out_of_range() {
495        let s = seg(0.0, 0.0, 1.0, 1.0);
496        assert!(s.eval(-0.1).is_err());
497        assert!(s.eval(1.1).is_err());
498    }
499
500    #[test]
501    fn test_incompatible_segments() {
502        let s1 = seg(0.0, 0.0, 1.0, 1.0);
503        let s2 = seg(0.5, 0.0, 1.0, 1.0);
504        assert!(s1.above(&s2).is_err());
505        assert!(s1.above_eq(&s2).is_err());
506        assert!(s1.intersects(&s2).is_err());
507        assert!(s1.intersection(&s2).is_err());
508    }
509
510    #[test]
511    fn test_approximately_equal() {
512        let s1 = seg(0.0, 1.0, 1.0, 2.0);
513        let s2 = seg(0.0, 1.000_001, 1.0, 2.000_001);
514        let s3 = seg(0.0, 1.1, 1.0, 2.0);
515        assert!(s1.approximately_equal(&s2));
516        assert!(!s1.approximately_equal(&s3));
517        let z = seg(0.0, 0.0, 1.0, 0.0);
518        assert!(z.approximately_equal(&z));
519        assert!(!z.approximately_equal(&seg(0.0, 1e-9, 1.0, 0.0)));
520    }
521
522    #[test]
523    fn test_max_of_segments_empty() {
524        assert!(max_of_segments(&[]).is_err());
525    }
526
527    #[test]
528    fn test_max_of_segments_incompatible() {
529        let s1 = seg(0.0, 0.0, 1.0, 1.0);
530        let s2 = seg(0.5, 0.0, 1.0, 1.0);
531        assert!(max_of_segments(&[s1, s2]).is_err());
532    }
533}