use {
super::num::Num,
docext::docext,
std::{fmt, marker::PhantomData, ops},
};
#[docext]
pub trait Curve: Sized {
const SIZE: usize;
const P: Num;
#[docext]
const N: Num;
#[docext]
const A: Num;
#[docext]
const B: Num;
fn g() -> Point<Self>;
}
#[derive(Debug)]
pub struct Point<C>(Coordinates, PhantomData<C>);
impl<C> Clone for Point<C> {
fn clone(&self) -> Self {
*self
}
}
impl<C> Copy for Point<C> {}
impl<C> PartialEq for Point<C> {
fn eq(&self, other: &Self) -> bool {
self.0 == other.0
}
}
impl<C> Eq for Point<C> {}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[docext]
pub enum Coordinates {
Infinity,
Finite(Num, Num),
}
#[docext]
impl<C: Curve> ops::Add for Point<C> {
type Output = Self;
fn add(self, rhs: Self) -> Self::Output {
match (self.0, rhs.0) {
(Coordinates::Infinity, other) | (other, Coordinates::Infinity) => {
Self(other, Default::default())
}
(Coordinates::Finite(x1, y1), Coordinates::Finite(x2, y2)) if x1 == x2 && y1 == y2 => {
let Some(inv) = Num::TWO.mul(y1, C::P).inv(C::P) else {
return Self(Coordinates::Infinity, Default::default());
};
let h = Num::THREE.mul(x1, C::P).mul(x1, C::P).mul(inv, C::P);
let x = h.mul(h, C::P).sub(Num::TWO.mul(x1, C::P), C::P);
let s = x1.sub(x, C::P);
Self::new(x, h.mul(s, C::P).sub(y1, C::P)).unwrap()
}
(Coordinates::Finite(x1, y1), Coordinates::Finite(x2, y2)) => {
let Some(inv) = x2.sub(x1, C::P).inv(C::P) else {
return Self(Coordinates::Infinity, Default::default());
};
let h = y2.sub(y1, C::P).mul(inv, C::P);
let x = h.mul(h, C::P).sub(x1, C::P).sub(x2, C::P);
let s = x1.sub(x, C::P);
Self::new(x, h.mul(s, C::P).sub(y1, C::P)).unwrap()
}
}
}
}
impl<C: Curve> ops::AddAssign for Point<C> {
fn add_assign(&mut self, rhs: Self) {
*self = *self + rhs;
}
}
impl<C: Curve> Point<C> {
pub fn new(x: Num, y: Num) -> Result<Self, InvalidPoint> {
let y2 = y.mul(y, C::P);
let x3 = x.mul(x, C::P).mul(x, C::P);
let ax = C::A.mul(x, C::P);
if y2 == x3.add(ax, C::P).add(C::B, C::P) {
Ok(Self(Coordinates::Finite(x, y), Default::default()))
} else {
Err(InvalidPoint)
}
}
pub fn infinity() -> Self {
Self(Coordinates::Infinity, Default::default())
}
pub fn coordinates(&self) -> Coordinates {
self.0
}
pub(super) fn scale(&self, n: Num) -> Self {
let mut s = *self;
let mut result = Self::infinity();
for i in 0..Num::BITS {
if n.get_bit(i) {
result += s;
}
s += s;
}
result
}
}
#[derive(Debug, Clone, Copy)]
pub struct InvalidPoint;
impl fmt::Display for InvalidPoint {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "invalid point")
}
}
impl std::error::Error for InvalidPoint {}