Skip to main content

proof_engine/number_theory/
elliptic.rs

1//! Elliptic curves over the reals: point arithmetic, rendering, and group law.
2
3use glam::{Vec2, Vec3, Vec4};
4
5/// An elliptic curve in short Weierstrass form: y^2 = x^3 + ax + b.
6#[derive(Debug, Clone, Copy)]
7pub struct EllipticCurve {
8    pub a: f64,
9    pub b: f64,
10}
11
12/// A point on an elliptic curve (or the point at infinity).
13#[derive(Debug, Clone, Copy, PartialEq)]
14pub enum CurvePoint {
15    Infinity,
16    Point(f64, f64),
17}
18
19impl EllipticCurve {
20    pub fn new(a: f64, b: f64) -> Self {
21        Self { a, b }
22    }
23
24    /// Discriminant: -16(4a^3 + 27b^2). Non-zero means the curve is non-singular.
25    pub fn discriminant(&self) -> f64 {
26        -16.0 * (4.0 * self.a.powi(3) + 27.0 * self.b.powi(2))
27    }
28
29    /// Check whether the curve is non-singular.
30    pub fn is_non_singular(&self) -> bool {
31        self.discriminant().abs() > 1e-12
32    }
33
34    /// Evaluate the right-hand side: x^3 + ax + b.
35    pub fn rhs(&self, x: f64) -> f64 {
36        x * x * x + self.a * x + self.b
37    }
38
39    /// Check if a point lies on the curve (within tolerance).
40    pub fn is_on_curve(&self, point: CurvePoint) -> bool {
41        match point {
42            CurvePoint::Infinity => true,
43            CurvePoint::Point(x, y) => {
44                let lhs = y * y;
45                let rhs = self.rhs(x);
46                (lhs - rhs).abs() < 1e-8
47            }
48        }
49    }
50
51    /// Add two points on the curve (elliptic curve group law).
52    pub fn point_add(&self, p: CurvePoint, q: CurvePoint) -> CurvePoint {
53        match (p, q) {
54            (CurvePoint::Infinity, _) => q,
55            (_, CurvePoint::Infinity) => p,
56            (CurvePoint::Point(x1, y1), CurvePoint::Point(x2, y2)) => {
57                // Check if P = -Q (same x, opposite y)
58                if (x1 - x2).abs() < 1e-12 && (y1 + y2).abs() < 1e-12 {
59                    return CurvePoint::Infinity;
60                }
61
62                let m = if (x1 - x2).abs() < 1e-12 {
63                    // Point doubling: m = (3x1^2 + a) / (2y1)
64                    if y1.abs() < 1e-12 {
65                        return CurvePoint::Infinity;
66                    }
67                    (3.0 * x1 * x1 + self.a) / (2.0 * y1)
68                } else {
69                    // General addition: m = (y2 - y1) / (x2 - x1)
70                    (y2 - y1) / (x2 - x1)
71                };
72
73                let x3 = m * m - x1 - x2;
74                let y3 = m * (x1 - x3) - y1;
75                CurvePoint::Point(x3, y3)
76            }
77        }
78    }
79
80    /// Scalar multiplication: compute n * P using double-and-add.
81    pub fn scalar_multiply(&self, n: i64, p: CurvePoint) -> CurvePoint {
82        if n == 0 {
83            return CurvePoint::Infinity;
84        }
85        let (mut k, point) = if n < 0 {
86            (-n as u64, self.negate(p))
87        } else {
88            (n as u64, p)
89        };
90
91        let mut result = CurvePoint::Infinity;
92        let mut base = point;
93        while k > 0 {
94            if k & 1 == 1 {
95                result = self.point_add(result, base);
96            }
97            base = self.point_add(base, base);
98            k >>= 1;
99        }
100        result
101    }
102
103    /// Negate a point: (x, y) -> (x, -y).
104    pub fn negate(&self, p: CurvePoint) -> CurvePoint {
105        match p {
106            CurvePoint::Infinity => CurvePoint::Infinity,
107            CurvePoint::Point(x, y) => CurvePoint::Point(x, -y),
108        }
109    }
110
111    /// Sample points on the curve for rendering: returns both upper and lower branches.
112    pub fn sample_curve(&self, x_min: f64, x_max: f64, steps: usize) -> Vec<Vec2> {
113        let mut points = Vec::new();
114        let dx = (x_max - x_min) / steps as f64;
115        for i in 0..=steps {
116            let x = x_min + i as f64 * dx;
117            let rhs = self.rhs(x);
118            if rhs >= 0.0 {
119                let y = rhs.sqrt();
120                points.push(Vec2::new(x as f32, y as f32));
121                points.push(Vec2::new(x as f32, -y as f32));
122            }
123        }
124        points
125    }
126
127    /// Find a point on the curve near x by solving y^2 = rhs.
128    pub fn point_at_x(&self, x: f64) -> Option<CurvePoint> {
129        let rhs = self.rhs(x);
130        if rhs < 0.0 {
131            None
132        } else {
133            Some(CurvePoint::Point(x, rhs.sqrt()))
134        }
135    }
136}
137
138// ─── Renderer ───────────────────────────────────────────────────────────────
139
140/// Renders the elliptic curve and point operations as glyph paths.
141pub struct EllipticCurveRenderer {
142    pub curve: EllipticCurve,
143    pub origin: Vec3,
144    pub scale: f32,
145    pub x_range: (f64, f64),
146}
147
148pub struct CurveGlyph {
149    pub position: Vec3,
150    pub color: Vec4,
151    pub character: char,
152}
153
154impl EllipticCurveRenderer {
155    pub fn new(curve: EllipticCurve, origin: Vec3, scale: f32, x_range: (f64, f64)) -> Self {
156        Self { curve, origin, scale, x_range }
157    }
158
159    /// Render the curve itself as a series of glyphs.
160    pub fn render_curve(&self, steps: usize) -> Vec<CurveGlyph> {
161        let points = self.curve.sample_curve(self.x_range.0, self.x_range.1, steps);
162        points
163            .iter()
164            .map(|p| {
165                let pos = self.origin + Vec3::new(p.x * self.scale, p.y * self.scale, 0.0);
166                let t = ((p.x - self.x_range.0 as f32)
167                    / (self.x_range.1 - self.x_range.0) as f32)
168                    .clamp(0.0, 1.0);
169                CurveGlyph {
170                    position: pos,
171                    color: Vec4::new(0.2, 0.8, t, 1.0),
172                    character: '.',
173                }
174            })
175            .collect()
176    }
177
178    /// Render a point addition: P, Q, and P+Q with connecting lines.
179    pub fn render_addition(
180        &self,
181        p: CurvePoint,
182        q: CurvePoint,
183    ) -> Vec<CurveGlyph> {
184        let mut glyphs = Vec::new();
185
186        let pq = self.curve.point_add(p, q);
187
188        let point_data = [
189            (p, Vec4::new(1.0, 0.3, 0.3, 1.0), 'P'),
190            (q, Vec4::new(0.3, 1.0, 0.3, 1.0), 'Q'),
191            (pq, Vec4::new(0.3, 0.3, 1.0, 1.0), 'R'),
192        ];
193
194        for &(pt, color, ch) in &point_data {
195            if let CurvePoint::Point(x, y) = pt {
196                glyphs.push(CurveGlyph {
197                    position: self.origin
198                        + Vec3::new(x as f32 * self.scale, y as f32 * self.scale, 0.0),
199                    color,
200                    character: ch,
201                });
202            }
203        }
204
205        // Add line from P to Q
206        if let (CurvePoint::Point(x1, y1), CurvePoint::Point(x2, y2)) = (p, q) {
207            let line_steps = 20;
208            for i in 0..=line_steps {
209                let t = i as f64 / line_steps as f64;
210                let x = x1 + t * (x2 - x1);
211                let y = y1 + t * (y2 - y1);
212                glyphs.push(CurveGlyph {
213                    position: self.origin
214                        + Vec3::new(x as f32 * self.scale, y as f32 * self.scale, 0.0),
215                    color: Vec4::new(0.5, 0.5, 0.5, 0.5),
216                    character: '-',
217                });
218            }
219        }
220
221        glyphs
222    }
223
224    /// Render scalar multiples: P, 2P, 3P, ..., nP.
225    pub fn render_multiples(&self, p: CurvePoint, n: i64) -> Vec<CurveGlyph> {
226        let mut glyphs = Vec::new();
227        for k in 1..=n {
228            let kp = self.curve.scalar_multiply(k, p);
229            if let CurvePoint::Point(x, y) = kp {
230                let t = k as f32 / n as f32;
231                glyphs.push(CurveGlyph {
232                    position: self.origin
233                        + Vec3::new(x as f32 * self.scale, y as f32 * self.scale, 0.0),
234                    color: Vec4::new(t, 0.5, 1.0 - t, 1.0),
235                    character: std::char::from_digit(k as u32 % 10, 10).unwrap_or('*'),
236                });
237            }
238        }
239        glyphs
240    }
241}
242
243// ─── Tests ──────────────────────────────────────────────────────────────────
244
245#[cfg(test)]
246mod tests {
247    use super::*;
248
249    fn approx(a: f64, b: f64, eps: f64) -> bool {
250        (a - b).abs() < eps
251    }
252
253    fn point_approx_eq(p: CurvePoint, q: CurvePoint, eps: f64) -> bool {
254        match (p, q) {
255            (CurvePoint::Infinity, CurvePoint::Infinity) => true,
256            (CurvePoint::Point(x1, y1), CurvePoint::Point(x2, y2)) => {
257                approx(x1, x2, eps) && approx(y1, y2, eps)
258            }
259            _ => false,
260        }
261    }
262
263    #[test]
264    fn discriminant() {
265        // y^2 = x^3 - x (a=-1, b=0): disc = -16(4*(-1)^3 + 0) = -16*(-4) = 64
266        let e = EllipticCurve::new(-1.0, 0.0);
267        assert!(approx(e.discriminant(), 64.0, 1e-10));
268        assert!(e.is_non_singular());
269    }
270
271    #[test]
272    fn on_curve() {
273        let e = EllipticCurve::new(-1.0, 1.0); // y^2 = x^3 - x + 1
274        // x=0: y^2 = 1, y = 1
275        assert!(e.is_on_curve(CurvePoint::Point(0.0, 1.0)));
276        assert!(e.is_on_curve(CurvePoint::Point(0.0, -1.0)));
277        assert!(e.is_on_curve(CurvePoint::Infinity));
278        assert!(!e.is_on_curve(CurvePoint::Point(0.0, 0.5)));
279    }
280
281    #[test]
282    fn identity_element() {
283        let e = EllipticCurve::new(-1.0, 1.0);
284        let p = CurvePoint::Point(0.0, 1.0);
285        // P + O = P
286        assert!(point_approx_eq(e.point_add(p, CurvePoint::Infinity), p, 1e-10));
287        // O + P = P
288        assert!(point_approx_eq(e.point_add(CurvePoint::Infinity, p), p, 1e-10));
289    }
290
291    #[test]
292    fn inverse_element() {
293        let e = EllipticCurve::new(-1.0, 1.0);
294        let p = CurvePoint::Point(0.0, 1.0);
295        let neg_p = e.negate(p);
296        // P + (-P) = O
297        assert_eq!(e.point_add(p, neg_p), CurvePoint::Infinity);
298    }
299
300    #[test]
301    fn point_addition() {
302        let e = EllipticCurve::new(-1.0, 1.0); // y^2 = x^3 - x + 1
303        let p = CurvePoint::Point(0.0, 1.0);
304        let q = CurvePoint::Point(1.0, 1.0);
305        let r = e.point_add(p, q);
306        // Verify result is on curve
307        assert!(e.is_on_curve(r), "P+Q should be on the curve");
308    }
309
310    #[test]
311    fn point_doubling() {
312        let e = EllipticCurve::new(-1.0, 1.0);
313        let p = CurvePoint::Point(0.0, 1.0);
314        let twop = e.point_add(p, p);
315        assert!(e.is_on_curve(twop), "2P should be on the curve");
316    }
317
318    #[test]
319    fn associativity() {
320        let e = EllipticCurve::new(-1.0, 1.0);
321        let p = CurvePoint::Point(0.0, 1.0);
322        let q = CurvePoint::Point(1.0, 1.0);
323        // (P+Q)+P vs P+(Q+P)
324        let lhs = e.point_add(e.point_add(p, q), p);
325        let rhs = e.point_add(p, e.point_add(q, p));
326        assert!(
327            point_approx_eq(lhs, rhs, 1e-8),
328            "associativity failed: {:?} vs {:?}",
329            lhs,
330            rhs
331        );
332    }
333
334    #[test]
335    fn commutativity() {
336        let e = EllipticCurve::new(-1.0, 1.0);
337        let p = CurvePoint::Point(0.0, 1.0);
338        let q = CurvePoint::Point(1.0, 1.0);
339        let pq = e.point_add(p, q);
340        let qp = e.point_add(q, p);
341        assert!(point_approx_eq(pq, qp, 1e-10));
342    }
343
344    #[test]
345    fn scalar_multiply_test() {
346        let e = EllipticCurve::new(-1.0, 1.0);
347        let p = CurvePoint::Point(0.0, 1.0);
348        // 0*P = O
349        assert_eq!(e.scalar_multiply(0, p), CurvePoint::Infinity);
350        // 1*P = P
351        assert!(point_approx_eq(e.scalar_multiply(1, p), p, 1e-10));
352        // 2*P = P+P
353        assert!(point_approx_eq(
354            e.scalar_multiply(2, p),
355            e.point_add(p, p),
356            1e-10
357        ));
358        // 3*P = P+P+P
359        assert!(point_approx_eq(
360            e.scalar_multiply(3, p),
361            e.point_add(e.point_add(p, p), p),
362            1e-8
363        ));
364    }
365
366    #[test]
367    fn scalar_multiply_negative() {
368        let e = EllipticCurve::new(-1.0, 1.0);
369        let p = CurvePoint::Point(0.0, 1.0);
370        let neg_p = e.negate(p);
371        assert!(point_approx_eq(e.scalar_multiply(-1, p), neg_p, 1e-10));
372        // P + (-P) = O
373        let sum = e.point_add(e.scalar_multiply(3, p), e.scalar_multiply(-3, p));
374        assert_eq!(sum, CurvePoint::Infinity);
375    }
376
377    #[test]
378    fn sample_curve_nonempty() {
379        let e = EllipticCurve::new(-1.0, 1.0);
380        let pts = e.sample_curve(-2.0, 3.0, 100);
381        assert!(!pts.is_empty());
382    }
383
384    #[test]
385    fn renderer_curve() {
386        let e = EllipticCurve::new(-1.0, 1.0);
387        let r = EllipticCurveRenderer::new(e, Vec3::ZERO, 1.0, (-2.0, 3.0));
388        let glyphs = r.render_curve(50);
389        assert!(!glyphs.is_empty());
390    }
391
392    #[test]
393    fn renderer_addition() {
394        let e = EllipticCurve::new(-1.0, 1.0);
395        let r = EllipticCurveRenderer::new(e, Vec3::ZERO, 1.0, (-2.0, 3.0));
396        let p = CurvePoint::Point(0.0, 1.0);
397        let q = CurvePoint::Point(1.0, 1.0);
398        let glyphs = r.render_addition(p, q);
399        assert!(glyphs.len() >= 3); // at least P, Q, R
400    }
401
402    #[test]
403    fn singular_curve() {
404        // y^2 = x^3 (a=0, b=0), disc = 0
405        let e = EllipticCurve::new(0.0, 0.0);
406        assert!(!e.is_non_singular());
407    }
408}