| ... | @@ -12,7 +12,7 @@ pub const Ristretto255 = struct { | ... | @@ -12,7 +12,7 @@ pub const Ristretto255 = struct { |
| 12 | | 12 | |
| 13 | p: Curve, | 13 | p: Curve, |
| 14 | | 14 | |
| 15 | fn sqrtRatioM1(u: Fe, v: Fe) !Fe { | 15 | fn sqrtRatioM1(u: Fe, v: Fe) struct { ratio_is_square: u32, root: Fe } { |
| 16 | const v3 = v.sq().mul(v); // v^3 | 16 | const v3 = v.sq().mul(v); // v^3 |
| 17 | var x = v3.sq().mul(u).mul(v).pow2523().mul(v3).mul(u); // uv^3(uv^7)^((q-5)/8) | 17 | var x = v3.sq().mul(u).mul(v).pow2523().mul(v3).mul(u); // uv^3(uv^7)^((q-5)/8) |
| 18 | const vxx = x.sq().mul(v); // vx^2 | 18 | const vxx = x.sq().mul(v); // vx^2 |
| ... | @@ -24,11 +24,7 @@ pub const Ristretto255 = struct { | ... | @@ -24,11 +24,7 @@ pub const Ristretto255 = struct { |
| 24 | const has_f_root = f_root_check.isZero(); | 24 | const has_f_root = f_root_check.isZero(); |
| 25 | const x_sqrtm1 = x.mul(Fe.sqrtm1); // x*sqrt(-1) | 25 | const x_sqrtm1 = x.mul(Fe.sqrtm1); // x*sqrt(-1) |
| 26 | x.cMov(x_sqrtm1, @boolToInt(has_p_root) | @boolToInt(has_f_root)); | 26 | x.cMov(x_sqrtm1, @boolToInt(has_p_root) | @boolToInt(has_f_root)); |
| 27 | const xa = x.abs(); | 27 | return .{ .ratio_is_square = @boolToInt(has_m_root) | @boolToInt(has_p_root), .root = x.abs() }; |
| 28 | if ((@boolToInt(has_m_root) | @boolToInt(has_p_root)) == 0) { | | |
| 29 | return error.NoRoot; | | |
| 30 | } | | |
| 31 | return xa; | | |
| 32 | } | 28 | } |
| 33 | | 29 | |
| 34 | fn rejectNonCanonical(s: [32]u8) !void { | 30 | fn rejectNonCanonical(s: [32]u8) !void { |
| ... | @@ -57,15 +53,14 @@ pub const Ristretto255 = struct { | ... | @@ -57,15 +53,14 @@ pub const Ristretto255 = struct { |
| 57 | const u2u2 = u2_.sq(); // (1+s^2)^2 | 53 | const u2u2 = u2_.sq(); // (1+s^2)^2 |
| 58 | const v = Fe.edwards25519d.mul(u1u1).neg().sub(u2u2); // -(d*u1^2)-u2^2 | 54 | const v = Fe.edwards25519d.mul(u1u1).neg().sub(u2u2); // -(d*u1^2)-u2^2 |
| 59 | const v_u2u2 = v.mul(u2u2); // v*u2^2 | 55 | const v_u2u2 = v.mul(u2u2); // v*u2^2 |
| 60 | const inv_sqrt = sqrtRatioM1(Fe.one, v_u2u2) catch |e| { | 56 | |
| 61 | return error.InvalidEncoding; | 57 | const inv_sqrt = sqrtRatioM1(Fe.one, v_u2u2); |
| 62 | }; | 58 | var x = inv_sqrt.root.mul(u2_); |
| 63 | var x = inv_sqrt.mul(u2_); | 59 | const y = inv_sqrt.root.mul(x).mul(v).mul(u1_); |
| 64 | const y = inv_sqrt.mul(x).mul(v).mul(u1_); | | |
| 65 | x = x.mul(s_); | 60 | x = x.mul(s_); |
| 66 | x = x.add(x).abs(); | 61 | x = x.add(x).abs(); |
| 67 | const t = x.mul(y); | 62 | const t = x.mul(y); |
| 68 | if ((@boolToInt(t.isNegative()) | @boolToInt(y.isZero())) != 0) { | 63 | if ((1 - inv_sqrt.ratio_is_square) | @boolToInt(t.isNegative()) | @boolToInt(y.isZero()) != 0) { |
| 69 | return error.InvalidEncoding; | 64 | return error.InvalidEncoding; |
| 70 | } | 65 | } |
| 71 | const p: Curve = .{ | 66 | const p: Curve = .{ |
| ... | @@ -85,9 +80,9 @@ pub const Ristretto255 = struct { | ... | @@ -85,9 +80,9 @@ pub const Ristretto255 = struct { |
| 85 | u1_ = u1_.mul(zmy); // (Z+Y)*(Z-Y) | 80 | u1_ = u1_.mul(zmy); // (Z+Y)*(Z-Y) |
| 86 | const u2_ = p.x.mul(p.y); // X*Y | 81 | const u2_ = p.x.mul(p.y); // X*Y |
| 87 | const u1_u2u2 = u2_.sq().mul(u1_); // u1*u2^2 | 82 | const u1_u2u2 = u2_.sq().mul(u1_); // u1*u2^2 |
| 88 | const inv_sqrt = sqrtRatioM1(Fe.one, u1_u2u2) catch unreachable; | 83 | const inv_sqrt = sqrtRatioM1(Fe.one, u1_u2u2); |
| 89 | const den1 = inv_sqrt.mul(u1_); | 84 | const den1 = inv_sqrt.root.mul(u1_); |
| 90 | const den2 = inv_sqrt.mul(u2_); | 85 | const den2 = inv_sqrt.root.mul(u2_); |
| 91 | const z_inv = den1.mul(den2).mul(p.t); // den1*den2*T | 86 | const z_inv = den1.mul(den2).mul(p.t); // den1*den2*T |
| 92 | const ix = p.x.mul(Fe.sqrtm1); // X*sqrt(-1) | 87 | const ix = p.x.mul(Fe.sqrtm1); // X*sqrt(-1) |
| 93 | const iy = p.y.mul(Fe.sqrtm1); // Y*sqrt(-1) | 88 | const iy = p.y.mul(Fe.sqrtm1); // Y*sqrt(-1) |
| ... | @@ -109,6 +104,35 @@ pub const Ristretto255 = struct { | ... | @@ -109,6 +104,35 @@ pub const Ristretto255 = struct { |
| 109 | return p.z.sub(y).mul(den_inv).abs().toBytes(); | 104 | return p.z.sub(y).mul(den_inv).abs().toBytes(); |
| 110 | } | 105 | } |
| 111 | | 106 | |
| | 107 | fn elligator(t: Fe) Curve { |
| | 108 | const r = t.sq().mul(Fe.sqrtm1); // sqrt(-1)*t^2 |
| | 109 | const u = r.add(Fe.one).mul(Fe.edwards25519eonemsqd); // (r+1)*(1-d^2) |
| | 110 | var c = comptime Fe.one.neg(); // -1 |
| | 111 | const v = c.sub(r.mul(Fe.edwards25519d)).mul(r.add(Fe.edwards25519d)); // (c-r*d)*(r+d) |
| | 112 | const ratio_sqrt = sqrtRatioM1(u, v); |
| | 113 | const wasnt_square = 1 - ratio_sqrt.ratio_is_square; |
| | 114 | var s = ratio_sqrt.root; |
| | 115 | const s_prime = s.mul(t).abs().neg(); // -|s*t| |
| | 116 | s.cMov(s_prime, wasnt_square); |
| | 117 | c.cMov(r, wasnt_square); |
| | 118 | |
| | 119 | const n = r.sub(Fe.one).mul(c).mul(Fe.edwards25519sqdmone).sub(v); // c*(r-1)*(d-1)^2-v |
| | 120 | const w0 = s.add(s).mul(v); // 2s*v |
| | 121 | const w1 = n.mul(Fe.edwards25519sqrtadm1); // n*sqrt(ad-1) |
| | 122 | const ss = s.sq(); // s^2 |
| | 123 | const w2 = Fe.one.sub(ss); // 1-s^2 |
| | 124 | const w3 = Fe.one.add(ss); // 1+s^2 |
| | 125 | |
| | 126 | return .{ .x = w0.mul(w3), .y = w2.mul(w1), .z = w1.mul(w3), .t = w0.mul(w2) }; |
| | 127 | } |
| | 128 | |
| | 129 | /// Map a 64-bit string into a Ristretto255 group element |
| | 130 | pub fn fromUniform(h: [64]u8) Ristretto255 { |
| | 131 | const p0 = elligator(Fe.fromBytes(h[0..32].*)); |
| | 132 | const p1 = elligator(Fe.fromBytes(h[32..64].*)); |
| | 133 | return Ristretto255{ .p = p0.add(p1) }; |
| | 134 | } |
| | 135 | |
| 112 | /// Double a Ristretto255 element. | 136 | /// Double a Ristretto255 element. |
| 113 | pub inline fn dbl(p: Ristretto255) Ristretto255 { | 137 | pub inline fn dbl(p: Ristretto255) Ristretto255 { |
| 114 | return .{ .p = p.p.dbl() }; | 138 | return .{ .p = p.p.dbl() }; |
| ... | @@ -125,6 +149,15 @@ pub const Ristretto255 = struct { | ... | @@ -125,6 +149,15 @@ pub const Ristretto255 = struct { |
| 125 | pub inline fn mul(p: Ristretto255, s: [32]u8) !Ristretto255 { | 149 | pub inline fn mul(p: Ristretto255, s: [32]u8) !Ristretto255 { |
| 126 | return Ristretto255{ .p = try p.p.mul(s) }; | 150 | return Ristretto255{ .p = try p.p.mul(s) }; |
| 127 | } | 151 | } |
| | 152 | |
| | 153 | /// Return true if two Ristretto255 elements are equivalent |
| | 154 | pub fn equivalent(p: Ristretto255, q: Ristretto255) bool { |
| | 155 | const p_ = &p.p; |
| | 156 | const q_ = &q.p; |
| | 157 | const a = p_.x.mul(q_.y).equivalent(p_.y.mul(q_.x)); |
| | 158 | const b = p_.y.mul(q_.y).equivalent(p_.x.mul(q_.x)); |
| | 159 | return (@boolToInt(a) | @boolToInt(b)) != 0; |
| | 160 | } |
| 128 | }; | 161 | }; |
| 129 | | 162 | |
| 130 | test "ristretto255" { | 163 | test "ristretto255" { |
| ... | @@ -141,4 +174,10 @@ test "ristretto255" { | ... | @@ -141,4 +174,10 @@ test "ristretto255" { |
| 141 | const s = [_]u8{15} ++ [_]u8{0} ** 31; | 174 | const s = [_]u8{15} ++ [_]u8{0} ** 31; |
| 142 | const w = try p.mul(s); | 175 | const w = try p.mul(s); |
| 143 | std.testing.expectEqualStrings(try std.fmt.bufPrint(&buf, "{X}", .{w.toBytes()}), "E0C418F7C8D9C4CDD7395B93EA124F3AD99021BB681DFC3302A9D99A2E53E64E"); | 176 | std.testing.expectEqualStrings(try std.fmt.bufPrint(&buf, "{X}", .{w.toBytes()}), "E0C418F7C8D9C4CDD7395B93EA124F3AD99021BB681DFC3302A9D99A2E53E64E"); |
| | 177 | |
| | 178 | std.testing.expect(p.dbl().dbl().dbl().dbl().equivalent(w.add(p))); |
| | 179 | |
| | 180 | const h = [_]u8{69} ** 32 ++ [_]u8{42} ** 32; |
| | 181 | const ph = Ristretto255.fromUniform(h); |
| | 182 | std.testing.expectEqualStrings(try std.fmt.bufPrint(&buf, "{X}", .{ph.toBytes()}), "DCCA54E037A4311EFBEEF413ACD21D35276518970B7A61DC88F8587B493D5E19"); |
| 144 | } | 183 | } |