1use glam::{Vec2, Vec3, Vec4};
4
5#[derive(Debug, Clone, Copy)]
7pub struct EllipticCurve {
8 pub a: f64,
9 pub b: f64,
10}
11
12#[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 pub fn discriminant(&self) -> f64 {
26 -16.0 * (4.0 * self.a.powi(3) + 27.0 * self.b.powi(2))
27 }
28
29 pub fn is_non_singular(&self) -> bool {
31 self.discriminant().abs() > 1e-12
32 }
33
34 pub fn rhs(&self, x: f64) -> f64 {
36 x * x * x + self.a * x + self.b
37 }
38
39 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 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 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 if y1.abs() < 1e-12 {
65 return CurvePoint::Infinity;
66 }
67 (3.0 * x1 * x1 + self.a) / (2.0 * y1)
68 } else {
69 (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 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 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 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 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
138pub 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 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 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 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 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#[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 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); 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 assert!(point_approx_eq(e.point_add(p, CurvePoint::Infinity), p, 1e-10));
287 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 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); 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 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 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 assert_eq!(e.scalar_multiply(0, p), CurvePoint::Infinity);
350 assert!(point_approx_eq(e.scalar_multiply(1, p), p, 1e-10));
352 assert!(point_approx_eq(
354 e.scalar_multiply(2, p),
355 e.point_add(p, p),
356 1e-10
357 ));
358 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 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); }
401
402 #[test]
403 fn singular_curve() {
404 let e = EllipticCurve::new(0.0, 0.0);
406 assert!(!e.is_non_singular());
407 }
408}