use elliptic_curve::{
Field,
hazmat::FieldArithmetic,
subtle::{Choice, ConditionallySelectable, ConstantTimeEq},
};
use crate::{AffinePoint, PrimeCurveParams};
pub trait AffineOsswuMap<C: PrimeCurveParams + FieldArithmetic<FieldElement: OsswuMap>> {
fn osswu(u: &C::FieldElement) -> Self;
}
impl<C> AffineOsswuMap<C> for AffinePoint<C>
where
C: PrimeCurveParams + FieldArithmetic<FieldElement: OsswuMap>,
{
fn osswu(u: &<C as FieldArithmetic>::FieldElement) -> Self {
let (x, y) = u.osswu();
Self { x, y, infinity: 0 }
}
}
#[derive(Debug)]
pub struct OsswuMapParams<F>
where
F: Field,
{
pub c1: &'static [u64],
pub c2: F,
pub map_a: F,
pub map_b: F,
pub z: F,
}
pub trait Sgn0 {
fn sgn0(&self) -> Choice;
}
pub trait OsswuMap: Field + Sgn0 {
const PARAMS: OsswuMapParams<Self>;
fn sqrt_ratio_3mod4(u: Self, v: Self) -> (Choice, Self) {
let tv1 = v.square();
let tv2 = u * v;
let tv1 = tv1 * tv2;
let y1 = tv1.pow_vartime(Self::PARAMS.c1);
let y1 = y1 * tv2;
let y2 = y1 * Self::PARAMS.c2;
let tv3 = y1.square();
let tv3 = tv3 * v;
let is_qr = tv3.ct_eq(&u);
let y = ConditionallySelectable::conditional_select(&y2, &y1, is_qr);
(is_qr, y)
}
fn osswu(&self) -> (Self, Self) {
let tv1 = self.square();
let tv1 = Self::PARAMS.z * tv1;
let tv2 = tv1.square();
let tv2 = tv2 + tv1;
let tv3 = tv2 + Self::ONE;
let tv3 = Self::PARAMS.map_b * tv3;
let tv4 = ConditionallySelectable::conditional_select(
&Self::PARAMS.z,
&-tv2,
!Field::is_zero(&tv2),
);
let tv4 = Self::PARAMS.map_a * tv4;
let tv2 = tv3.square();
let tv6 = tv4.square();
let tv5 = Self::PARAMS.map_a * tv6;
let tv2 = tv2 + tv5;
let tv2 = tv2 * tv3;
let tv6 = tv6 * tv4;
let tv5 = Self::PARAMS.map_b * tv6;
let tv2 = tv2 + tv5;
let x = tv1 * tv3;
let (is_gx1_square, y1) = Self::sqrt_ratio_3mod4(tv2, tv6);
let y = tv1 * self;
let y = y * y1;
let x = ConditionallySelectable::conditional_select(&x, &tv3, is_gx1_square);
let y = ConditionallySelectable::conditional_select(&y, &y1, is_gx1_square);
let e1 = self.sgn0().ct_eq(&y.sgn0());
let y = ConditionallySelectable::conditional_select(&-y, &y, e1);
let x = x * tv4.invert().unwrap();
(x, y)
}
}