feature. See also
. The project being documented here (as the example) is the Zig library itself.
ml_dsa.PolyVec
fn PolyVec(comptime len: u8) type
File
Code
fn PolyVec(comptime len: u8) type {
return struct {
ps: [len]Poly,
const Self = @This();
const zero: Self = .{ .ps = @splat(.zero) };
fn map(v: Self, comptime op: fn (Poly) Poly) Self {
var ret: Self = undefined;
inline for (0..len) |i| {
ret.ps[i] = op(v.ps[i]);
}
return ret;
}
fn mapBinary(a: Self, b: Self, comptime op: fn (Poly, Poly) Poly) Self {
var ret: Self = undefined;
inline for (0..len) |i| {
ret.ps[i] = op(a.ps[i], b.ps[i]);
}
return ret;
}
fn mapBinaryPoly(v: Self, scalar: Poly, comptime op: fn (Poly, Poly) Poly) Self {
var ret: Self = undefined;
inline for (0..len) |i| {
ret.ps[i] = op(v.ps[i], scalar);
}
return ret;
}
fn add(a: Self, b: Self) Self {
return mapBinary(a, b, Poly.add);
}
fn sub(a: Self, b: Self) Self {
return mapBinary(a, b, Poly.sub);
}
fn ntt(v: Self) Self {
return map(v, Poly.ntt);
}
fn invNTT(v: Self) Self {
return map(v, Poly.invNTT);
}
fn normalize(v: Self) Self {
return map(v, Poly.normalize);
}
fn reduceLe2Q(v: Self) Self {
return map(v, Poly.reduceLe2Q);
}
fn normalizeAssumingLe2Q(v: Self) Self {
return map(v, Poly.normalizeAssumingLe2Q);
}
fn exceeds(v: Self, bound: u32) bool {
var result = false;
for (0..len) |i| {
result = result or v.ps[i].exceeds(bound);
}
return result;
}
fn power2Round(v: Self, t0_out: *Self) Self {
var t1: Self = undefined;
for (0..len) |i| {
const result = v.ps[i].power2RoundPoly();
t0_out.ps[i] = result.t0;
t1.ps[i] = result.t1;
}
return t1;
}
fn packWith(
v: Self,
buf: []u8,
comptime poly_size: usize,
comptime pack_fn: fn (Poly, []u8) void,
) void {
inline for (0..len) |i| {
const offset = i * poly_size;
pack_fn(v.ps[i], buf[offset..][0..poly_size]);
}
}
fn unpackWith(
comptime poly_size: usize,
comptime unpack_fn: fn ([]const u8) Poly,
buf: []const u8,
) Self {
var result: Self = undefined;
inline for (0..len) |i| {
const offset = i * poly_size;
result.ps[i] = unpack_fn(buf[offset..][0..poly_size]);
}
return result;
}
fn packT1(v: Self, buf: []u8) void {
const poly_size = (N * (Q_BITS - D)) / 8;
packWith(v, buf, poly_size, polyPackT1);
}
fn unpackT1(bytes: []const u8) Self {
const poly_size = (N * (Q_BITS - D)) / 8;
return unpackWith(poly_size, polyUnpackT1, bytes);
}
fn packT0(v: Self, buf: []u8) void {
const poly_size = (N * D) / 8;
packWith(v, buf, poly_size, polyPackT0);
}
fn unpackT0(buf: []const u8) Self {
const poly_size = (N * D) / 8;
return unpackWith(poly_size, polyUnpackT0, buf);
}
fn packLeqEta(v: Self, comptime eta: u8, buf: []u8) void {
const poly_size = if (eta == 2) 96 else 128;
const pack_fn = struct {
fn pack(p: Poly, b: []u8) void {
polyPackLeqEta(p, eta, b);
}
}.pack;
packWith(v, buf, poly_size, pack_fn);
}
fn unpackLeqEta(comptime eta: u8, buf: []const u8) Self {
const poly_size = if (eta == 2) 96 else 128;
const unpack_fn = struct {
fn unpack(b: []const u8) Poly {
return polyUnpackLeqEta(eta, b);
}
}.unpack;
return unpackWith(poly_size, unpack_fn, buf);
}
fn packLeGamma1(v: Self, comptime gamma1_bits: u8, buf: []u8) void {
const poly_size = ((gamma1_bits + 1) * N) / 8;
const pack_fn = struct {
fn pack(p: Poly, b: []u8) void {
polyPackLeGamma1(p, gamma1_bits, b);
}
}.pack;
packWith(v, buf, poly_size, pack_fn);
}
fn unpackLeGamma1(comptime gamma1_bits: u8, buf: []const u8) Self {
const poly_size = ((gamma1_bits + 1) * N) / 8;
const unpack_fn = struct {
fn unpack(b: []const u8) Poly {
return polyUnpackLeGamma1(gamma1_bits, b);
}
}.unpack;
return unpackWith(poly_size, unpack_fn, buf);
}
fn packW1(v: Self, comptime gamma1_bits: u8, buf: []u8) void {
const poly_size = (N * (Q_BITS - gamma1_bits)) / 8;
const pack_fn = struct {
fn pack(p: Poly, b: []u8) void {
polyPackW1(p, gamma1_bits, b);
}
}.pack;
packWith(v, buf, poly_size, pack_fn);
}
fn decomposeVec(v: Self, comptime gamma2: u32, w0_out: *Self) Self {
var w1: Self = undefined;
for (0..len) |i| {
for (0..N) |j| {
const r = decompose(v.ps[i].cs[j], gamma2);
w0_out.ps[i].cs[j] = r.a0_plus_q;
w1.ps[i].cs[j] = r.a1;
}
}
return w1;
}
fn makeHintVec(w0mcs2pct0: Self, w1: Self, comptime gamma2: u32) struct { hint: Self, pop: u32 } {
var hint: Self = undefined;
var pop: u32 = 0;
for (0..len) |i| {
const result = polyMakeHint(w0mcs2pct0.ps[i], w1.ps[i], gamma2);
hint.ps[i] = result.hint;
pop += result.count;
}
return .{ .hint = hint, .pop = pop };
}
fn useHint(v: Self, hint: Self, comptime gamma2: u32) Self {
var result: Self = undefined;
for (0..len) |i| {
result.ps[i] = polyUseHint(v.ps[i], hint.ps[i], gamma2);
}
return result;
}
fn mulBy2toD(v: Self) Self {
var result: Self = undefined;
for (0..len) |i| {
for (0..N) |j| {
result.ps[i].cs[j] = v.ps[i].cs[j] << D;
}
}
return result;
}
fn deriveUniformLeGamma1(comptime gamma1_bits: u8, seed: *const [64]u8, nonce: u16) Self {
var result: Self = undefined;
for (0..len) |i| {
result.ps[i] = expandMask(gamma1_bits, seed, nonce + @as(u16, @intCast(i)));
}
return result;
}
fn packHint(v: Self, comptime omega: u16, buf: []u8) bool {
var idx: usize = 0;
var count: u32 = 0;
for (0..len) |i| {
for (0..N) |j| {
if (v.ps[i].cs[j] != 0) {
count += 1;
}
}
}
if (count > omega) {
return false;
}
// First omega bytes: positions of set bits across all polynomials
// Last len bytes: boundary indices showing where each polynomial's hints end
for (0..len) |i| {
for (0..N) |j| {
if (v.ps[i].cs[j] != 0) {
buf[idx] = @intCast(j);
idx += 1;
}
}
buf[omega + i] = @intCast(idx);
}
while (idx < omega) : (idx += 1) {
buf[idx] = 0;
}
return true;
}
fn unpackHint(comptime omega: u16, buf: []const u8) ?Self {
var result: Self = .{ .ps = @splat(.zero) };
var prev_sop: u8 = 0;
for (0..len) |i| {
const sop = buf[omega + i];
if (sop < prev_sop or sop > omega) {
return null;
}
var j = prev_sop;
while (j < sop) : (j += 1) {
if (j > prev_sop and buf[j] <= buf[j - 1]) {
return null;
}
const pos = buf[j];
if (pos >= N) {
return null;
}
result.ps[i].cs[pos] = 1;
}
prev_sop = sop;
}
var j = prev_sop;
while (j < omega) : (j += 1) {
if (buf[j] != 0) {
return null;
}
}
return result;
}
};
}