iuna

iuna

iuna - experimental mainnet-candidate protocol
git clone https://getiuna.org/git/iuna.git
Log | Files | Refs | README | LICENSE

commit aa1223b10c7c925699c2b7aca64f886bdd29bec6
parent fb71f0f7fbfcfec715eda75e5a40eda0e9071b0d
Author: Joris Hartog <jorishartog@hotmail.com>
Date:   Fri, 21 Aug 2026 13:34:32 +0200

Optimize Rust VDF prover

Diffstat:
MCargo.lock | 172++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++-----------------
MCargo.toml | 8+++++++-
Mdeployment.sh | 3+++
Mdocs/protocol.md | 6++++--
Mdocs/security-review.md | 39+++++++++++++++++++++++++++++++++++----
Mfuzz/Cargo.lock | 108++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++---
Mfuzz/Cargo.toml | 6++++++
Mfuzz/README.md | 2++
Afuzz/corpus/vdf_proof/chia-42-300.hex | 1+
Afuzz/fuzz_targets/vdf_proof.rs | 47+++++++++++++++++++++++++++++++++++++++++++++++
Msrc/domain/ledger_ops.rs | 3++-
Dsrc/domain/vdf.rs | 397-------------------------------------------------------------------------------
Asrc/domain/vdf/arithmetic.rs | 286+++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
Asrc/domain/vdf/limb_arithmetic.rs | 570+++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
Asrc/domain/vdf/limbs/mod.rs | 2413+++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
Asrc/domain/vdf/mod.rs | 216+++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
Asrc/domain/vdf/prover.rs | 591+++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
Asrc/domain/vdf/reducer.rs | 186+++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
Asrc/domain/vdf/reference.rs | 37+++++++++++++++++++++++++++++++++++++
Asrc/domain/vdf/wesolowski.rs | 130+++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++++
20 files changed, 4776 insertions(+), 445 deletions(-)

diff --git a/Cargo.lock b/Cargo.lock @@ -157,6 +157,12 @@ dependencies = [ ] [[package]] +name = "bumpalo" +version = "3.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5d20789868f4b01b2f2caec9f5c4e0213b41e3e5702a50157d699ae31ced2fcb" + +[[package]] name = "bytes" version = "1.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -164,9 +170,9 @@ checksum = "fc652a48c352aef3ea3aed32080501cf3ef6ed5da78602a020c991775b0aff04" [[package]] name = "cc" -version = "1.3.0" +version = "1.4.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c89588d05638b5b4594a3348a2d6c20277e43a7f5c5202b05cc56888475a47b8" +checksum = "509591b7bcd67f4ef775afad7662703b4935daaa6ec0e5605cfb1090b32a2b6d" dependencies = [ "find-msvc-tools", "shlex", @@ -347,9 +353,9 @@ checksum = "28dea519a9695b9977216879a3ebfddf92f1c08c05d984f8996aecd6ecdc811d" [[package]] name = "find-msvc-tools" -version = "0.1.9" +version = "0.1.11" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" +checksum = "d45db016d36b838f563236e9193d0ee6ce38f3f68b6c94e914b4929c96bbb890" [[package]] name = "fnv" @@ -368,30 +374,30 @@ dependencies = [ [[package]] name = "futures-channel" -version = "0.3.33" +version = "0.3.34" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "262590f4fe6afeb0bc83be1daa64e52657fe185690a958af7f3ad0e92085c5ae" +checksum = "b1f9e3d69d39e4862ffed03ed071a76f9a13ba1d9109d355b0f0aa6b15e393c4" dependencies = [ "futures-core", ] [[package]] name = "futures-core" -version = "0.3.33" +version = "0.3.34" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "2cd50c473c80f6d7c3670a752354b8e569b1a7cbfdc0419ec88e5edad85e0dc7" +checksum = "92d699e522242e69e3003b94ecc1f960f3a5e015aa7c5d7486e65ad01dd94f5e" [[package]] name = "futures-task" -version = "0.3.33" +version = "0.3.34" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b231ed28831efb4a61a08580c4bc233ec56bc009f4cd8f52da2c3cb97df0c109" +checksum = "cd417de3d1d015fc3bfd2b1ea46dfc7bab72ef86f1cc7cc9c78e728b34a6d1fd" [[package]] name = "futures-util" -version = "0.3.33" +version = "0.3.34" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a77a90a256fce34da66415271e30f94ee91c57b04b8a2c042d9cf3220179deaa" +checksum = "0d50a92467f8ba5dd6e3ee5d4bd04d73ab2e4e1c44474a0674821dfce14b79bc" dependencies = [ "futures-core", "futures-task", @@ -481,9 +487,9 @@ dependencies = [ [[package]] name = "http" -version = "1.4.2" +version = "1.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6970f50e31d6fc17d3fa27329444bfa74e196cf62e95052a3f6fee181dba6425" +checksum = "918d3568bebf352712bc2ef3d46a8bcf1a75b373be6539de198e9105cbbf9ce0" dependencies = [ "bytes", "itoa", @@ -501,9 +507,9 @@ dependencies = [ [[package]] name = "http-body-util" -version = "0.1.4" +version = "0.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e9f41fd6a08e4d4ec69df65976da761afd5ad5e58a9d4acb46bd1c953a9e3ff2" +checksum = "23169fe34a5fbcdd3f3862e78fb9b6fccd5f02a6dc6f732547005d45631ce71c" dependencies = [ "bytes", "futures-core", @@ -584,7 +590,9 @@ dependencies = [ "chacha20poly1305", "ed25519-dalek", "getrandom 0.2.17", + "kyn-vdf", "num-bigint", + "num-integer", "num-traits", "pbkdf2", "proptest", @@ -598,10 +606,25 @@ dependencies = [ ] [[package]] +name = "kyn-vdf" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6f19091872d9bc1ebb51b146a4c7b2ce12c5de594c5f8c7cd495d51407e7ab60" +dependencies = [ + "num-bigint", + "num-integer", + "num-traits", + "sha2", + "thiserror", + "wasm", + "wasm-bindgen", +] + +[[package]] name = "libc" -version = "0.2.188" +version = "0.2.189" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "22053b6a34f84abc97f9129e61334f40174659a1b9bd18c970b83db6a9a6348b" +checksum = "3eaf3ede3fee6db1a4c2ee091bf8a8b4dccdc6d17f656fb07896ee72867612f2" [[package]] name = "libsqlite3-sys" @@ -666,9 +689,9 @@ dependencies = [ [[package]] name = "num-bigint" -version = "0.4.8" +version = "0.4.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c89e69e7e0f03bea5ef08013795c25018e101932225a656383bd384495ecc367" +checksum = "a5e44f723f1133c9deac646763579fdb3ac745e418f2a7af9cd0c431da1f20b9" dependencies = [ "num-integer", "num-traits", @@ -761,9 +784,9 @@ dependencies = [ [[package]] name = "pkg-config" -version = "0.3.33" +version = "0.3.34" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" +checksum = "f6b464fbc74e149a392436b17d523f769e057cb6877f6a5c4618bc6f11800548" [[package]] name = "poly1305" @@ -939,6 +962,12 @@ dependencies = [ ] [[package]] +name = "rustversion" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f" + +[[package]] name = "rusty-fork" version = "0.3.1" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -995,7 +1024,7 @@ checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348" dependencies = [ "proc-macro2", "quote", - "syn 3.0.2", + "syn 3.0.3", ] [[package]] @@ -1121,9 +1150,9 @@ dependencies = [ [[package]] name = "syn" -version = "3.0.2" +version = "3.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a207d6d6a2b7fc470b80443726053f18a2481b7e1eee970597051596567987a3" +checksum = "53e9bae58849f64dfa4f5d5ae372c8341f7305f82a3868709269343628b659a3" dependencies = [ "proc-macro2", "quote", @@ -1150,6 +1179,26 @@ dependencies = [ ] [[package]] +name = "thiserror" +version = "2.0.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ec86235f5fcc2a73650310756d2ac5b138a5780bbbdfae3eeccec992c435ba4f" +dependencies = [ + "thiserror-impl", +] + +[[package]] +name = "thiserror-impl" +version = "2.0.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bc04cd3e1236dd4a98afca4569f2deb3f120e5422a4023be2cb683f8486292af" +dependencies = [ + "proc-macro2", + "quote", + "syn 3.0.3", +] + +[[package]] name = "tinyvec" version = "1.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1183,13 +1232,13 @@ dependencies = [ [[package]] name = "tokio-macros" -version = "2.7.1" +version = "2.7.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6328af13490e73a9b4694030fafd93f8c8c6a9dede33e821c3fc63eddf8042ba" +checksum = "78773a2a397f451582ce068015985c33193cf6dea8b74d2a639fe457b2f07b0e" dependencies = [ "proc-macro2", "quote", - "syn 2.0.119", + "syn 3.0.3", ] [[package]] @@ -1306,14 +1355,65 @@ checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" [[package]] name = "wasip2" -version = "1.0.1+wasi-0.2.4" +version = "1.0.4+wasi-0.2.12" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0562428422c63773dad2c345a1882263bbf4d65cf3f42e90921f787ef5ad58e7" +checksum = "b67efb37e106e55ce722a510d6b5f9c17f083e5fc79afc2badeb12cc313d9487" dependencies = [ "wit-bindgen", ] [[package]] +name = "wasm" +version = "0.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e59fe57342dd136b22e8c7d1856e276b69e07a8f88153d218d2e54f19991a3eb" + +[[package]] +name = "wasm-bindgen" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ddb3f79143bced6de84270411622a2699cee572fc0875aeaf1e7867cf9fca1a" +dependencies = [ + "cfg-if", + "once_cell", + "rustversion", + "wasm-bindgen-macro", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-macro" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4e21a184b13fb19e157296e2c46056aec9092264fab83e4ba59e68c61b323c3d" +dependencies = [ + "quote", + "wasm-bindgen-macro-support", +] + +[[package]] +name = "wasm-bindgen-macro-support" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fecefd9c35bd935a20fc3fc344b5f29138961e4f47fb03297d88f2587afb5ebd" +dependencies = [ + "bumpalo", + "proc-macro2", + "quote", + "syn 2.0.119", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-shared" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "23939e44bb9a5d7576fa2b563dc2e136628f1224e88a8deed09e04858b77871f" +dependencies = [ + "unicode-ident", +] + +[[package]] name = "windows-link" version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1330,24 +1430,24 @@ dependencies = [ [[package]] name = "wit-bindgen" -version = "0.46.0" +version = "0.57.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f17a85883d4e6d00e8a97c586de764dabcc06133f7f1d55dce5cdc070ad7fe59" +checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" [[package]] name = "zerocopy" -version = "0.8.55" +version = "0.8.56" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b5a105cd7b140f6eeec8acff2ea38135d3cab283ada58540f629fe51e46696eb" +checksum = "556764e583adb45a9f8d413c2a147fa7e8d821e48e12b14fd560b607998b75eb" dependencies = [ "zerocopy-derive", ] [[package]] name = "zerocopy-derive" -version = "0.8.55" +version = "0.8.56" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0fe976fb70c78cd64cccfe3a6fc142244e8a77b70959b30faf9d0ac37ee228eb" +checksum = "f2ab42fc20575779bd240faa45f94a74256f755c0fa9e89f0ede20d91d0cdfc1" dependencies = [ "proc-macro2", "quote", diff --git a/Cargo.toml b/Cargo.toml @@ -11,7 +11,8 @@ bip39 = "2.2.2" chacha20poly1305 = "0.10.1" ed25519-dalek = "2.2.0" getrandom = "0.2.17" -num-bigint = "0.4.6" +num-bigint = "=0.4.6" +num-integer = "=0.1.46" num-traits = "0.2.19" pbkdf2 = "0.12.2" serde = { version = "1.0.228", features = ["derive"] } @@ -19,6 +20,7 @@ serde_json = "1.0.150" sha2 = "0.10.9" rusqlite = { version = "0.32.1", features = ["bundled"] } tokio = { version = "1.45.1", features = ["full"] } +kyn-vdf = "=0.1.1" [features] fuzzing = [] @@ -27,3 +29,7 @@ fuzzing = [] proptest = "1.6.0" tempfile = "3.20.0" tower = { version = "0.5.2", features = ["util"] } + +[profile.release] +codegen-units = 1 +lto = "thin" diff --git a/deployment.sh b/deployment.sh @@ -82,7 +82,9 @@ run_release_tests() { require_command cargo local fuzz_runs="${IUNA_FUZZ_RUNS:-256}" + local vdf_fuzz_runs="${IUNA_VDF_FUZZ_RUNS:-16}" validate_positive_integer IUNA_FUZZ_RUNS "$fuzz_runs" + validate_positive_integer IUNA_VDF_FUZZ_RUNS "$vdf_fuzz_runs" cargo test --locked cargo check --locked --manifest-path fuzz/Cargo.toml @@ -91,6 +93,7 @@ run_release_tests() { cargo run --locked --manifest-path fuzz/Cargo.toml --bin domain_json -- -runs="$fuzz_runs" fuzz/corpus/domain_json cargo run --locked --manifest-path fuzz/Cargo.toml --bin stratum_request -- -runs="$fuzz_runs" fuzz/corpus/stratum_request cargo run --locked --manifest-path fuzz/Cargo.toml --bin wallet_config -- -runs="$fuzz_runs" fuzz/corpus/wallet_config + cargo run --locked --manifest-path fuzz/Cargo.toml --bin vdf_proof -- -runs="$vdf_fuzz_runs" fuzz/corpus/vdf_proof cargo test --locked --release --test properties -- --ignored } diff --git a/docs/protocol.md b/docs/protocol.md @@ -90,9 +90,11 @@ The anchor burn is not a fairness mechanism. By itself, it would mostly help the ## VDF Timing -The VDF is there to make block production sequential and time-based. It uses repeated squaring in an unknown-order RSA group: validators know the public modulus, but not its factorization. That unknown order is essential. If the factorization were known, a finalizer could skip the delay with normal modular exponentiation. +The VDF is there to make block production sequential and time-based. It uses a Chia-compatible Wesolowski proof over a class group of imaginary quadratic forms. The 1024-bit class-group discriminant is derived deterministically from the block VDF seed, so the protocol does not rely on an RSA trusted setup or on anyone destroying hidden factors. -The devnet uses the public RSA-2048 challenge modulus. A production mainnet should use a purpose-specific trusted setup ceremony with destroyed factors, or a class-group VDF that avoids trusted setup. +VDF solutions are encoded with the `classgroup-wesolowski-bqfc-v1` prefix followed by two 100-byte Chia BQFC forms in hexadecimal: the output `y` and the Wesolowski proof `pi`. The implementation is Rust-only and has no GMP, MPIR, or other native runtime dependency. Proof generation uses a Chia-compatible checkpoint-and-bucket time-memory tradeoff, with a bounded-memory constant-space fallback for parameter sets that exceed the local allocation limits. Both paths produce the same proof and do not change verification or the wire format. Older RSA-modulus and GMP class-group VDF outputs are not valid for this protocol version. + +On the local Apple Silicon release benchmark for 100,000 rounds with seed `iuna-vdf-prover-benchmark`, the current Rust-only checkpoint prover completed in about `0.89s` to `0.92s`, and the constant-memory prover completed in about `1.78s`, after moving output squaring, proof composition, and proof squaring onto the local custom `Vec<u64>` limb backend with reusable division/GCD/reduction scratch buffers, Lehmer-style full and partial XGCD batching, x-only extended-GCD paths for call sites that do not need the second Bezout coefficient, positive-input left-GCD fast paths, mutable Lehmer linear-combination outputs for XGCD batch updates, `u64` Lehmer quotient windows, scratch-backed scalar combinations, one-limb scalar multiplication into scratch buffers, quotient/remainder-directed division outputs, exact power-of-two division fast paths, clone-free signed subtraction, scratch-backed reduction steps with small-quotient fast paths and quotient comparison that avoids temporary doubled limbs, tighter add/sub limb loops, one-limb multiplication, small-shift fast paths, sparse proof buckets that keep empty buckets implicit instead of cloning full identity forms or composing identity aggregates, a 100,000-round checkpoint parameter floor of `k = 10`, release thin-LTO/codegen-unit tuning, and per-pass incremental checkpoint bucket selection that replaces per-checkpoint modular exponentiation with one modular exponentiation plus fixed modular steps. A separate official Python/C++ `chiavdf.prove()` reference run completed the same workload in about `0.673s`. Phase profiling measured the checkpoint prover at about `0.78s` to `0.81s` for output squaring and about `0.11s` to `0.12s` for proof construction. The limb backend covers signed limb arithmetic, division, full and partial XGCD, production NUDUPL/NUCOMP, checkpoint bucket selection, and class-group exponentiation. Closing the remaining gap requires deeper in-place arithmetic for multiplication intermediates, stronger multiplication algorithms for larger operands, and Windows MSVC release benchmarking without introducing GMP or MPIR. The target block time is `10 minutes`. The protocol retargets VDF rounds from recent observed block times: diff --git a/docs/security-review.md b/docs/security-review.md @@ -159,10 +159,41 @@ see which revision was tested. ## Launch-Blocking Review Items -- VDF trust assumption: `docs/protocol.md` documents that the devnet uses the - public RSA-2048 challenge modulus and that production mainnet should use a - purpose-specific trusted setup or class-group VDF. Before promotion, decide - whether this is accepted for the candidate or blocks mainnet. +- VDF implementation: `docs/protocol.md` documents the Rust-only, + Chia-compatible class-group Wesolowski VDF and its + `classgroup-wesolowski-bqfc-v1` solution encoding. The implementation must not + depend on GMP, MPIR, or native runtime libraries. On local Apple Silicon, a + 100,000-round release benchmark with seed `iuna-vdf-prover-benchmark` measured + about `0.89s` to `0.92s` for the current Rust-only checkpoint prover and about `1.78s` + for the Rust-only constant-memory prover after moving output squaring, proof + composition, and proof squaring onto the local custom `Vec<u64>` limb backend + with reusable division/GCD/reduction scratch buffers, Lehmer-style full and + partial XGCD batching, x-only extended-GCD paths for call sites that do not + need the second Bezout coefficient, positive-input left-GCD fast paths, + mutable Lehmer linear-combination outputs for XGCD batch updates, `u64` + Lehmer quotient windows, scratch-backed scalar combinations, one-limb scalar + multiplication into scratch buffers, quotient/remainder-directed + division outputs, exact power-of-two division fast paths, and owned + add/sub/shift helpers for formula temporaries, clone-free signed subtraction, + scratch-backed reduction steps with small-quotient fast paths and quotient + comparison that avoids temporary doubled limbs, + tighter add/sub limb loops, one-limb multiplication, small-shift fast paths, + plus sparse proof buckets that keep empty buckets implicit instead of cloning + full identity forms or composing identity aggregates, and per-pass incremental + checkpoint bucket selection with a 100,000-round checkpoint parameter floor of + `k = 10`, release thin-LTO/codegen-unit tuning, and replacement of + per-checkpoint modular exponentiation with one modular exponentiation plus + fixed modular steps; an official + Python/C++ `chiavdf` reference measurement was about `0.673s`. Phase profiling + measured about `0.78s` to `0.81s` in output squaring and about `0.11s` to `0.12s` in proof + construction. The limb backend covers signed limb arithmetic, division, full + and partial XGCD, production NUDUPL/NUCOMP, checkpoint bucket selection, and + class-group exponentiation. + Before promotion, review the + limb backend's canonical encoding and division behavior, add deeper in-place + arithmetic for multiplication intermediates where profiling justifies it, and run + fixed vectors, differential VDF tests, fuzz targets, release benchmarks, and + Windows MSVC builds on every release platform. - Wallet/key handling: review the encrypted wallet format, PBKDF2 iteration count, password UX, recovery phrase exposure, and backup guidance. - Public exposure: verify bootnodes expose only the intended P2P and optional diff --git a/fuzz/Cargo.lock b/fuzz/Cargo.lock @@ -148,6 +148,12 @@ dependencies = [ ] [[package]] +name = "bumpalo" +version = "3.20.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5d20789868f4b01b2f2caec9f5c4e0213b41e3e5702a50157d699ae31ced2fcb" + +[[package]] name = "bytes" version = "1.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -553,7 +559,9 @@ dependencies = [ "chacha20poly1305", "ed25519-dalek", "getrandom 0.2.17", + "kyn-vdf", "num-bigint", + "num-integer", "num-traits", "pbkdf2", "rusqlite", @@ -583,6 +591,21 @@ dependencies = [ ] [[package]] +name = "kyn-vdf" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6f19091872d9bc1ebb51b146a4c7b2ce12c5de594c5f8c7cd495d51407e7ab60" +dependencies = [ + "num-bigint", + "num-integer", + "num-traits", + "sha2", + "thiserror", + "wasm", + "wasm-bindgen", +] + +[[package]] name = "libc" version = "0.2.189" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -655,9 +678,9 @@ dependencies = [ [[package]] name = "num-bigint" -version = "0.4.8" +version = "0.4.6" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c89e69e7e0f03bea5ef08013795c25018e101932225a656383bd384495ecc367" +checksum = "a5e44f723f1133c9deac646763579fdb3ac745e418f2a7af9cd0c431da1f20b9" dependencies = [ "num-integer", "num-traits", @@ -665,9 +688,9 @@ dependencies = [ [[package]] name = "num-integer" -version = "0.1.47" +version = "0.1.46" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7ce2d95d4b3734dc35aa2f45e1aa22cd416814592a4f9d9205e11affd5b8e10b" +checksum = "7969661fd2958a5cb096e56c8e1ad0444ac2bbcd0061bd28660485a44879858f" dependencies = [ "num-traits", ] @@ -831,6 +854,12 @@ dependencies = [ ] [[package]] +name = "rustversion" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f" + +[[package]] name = "ryu" version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1017,6 +1046,26 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0bf256ce5efdfa370213c1dabab5935a12e49f2c58d15e9eac2870d3b4f27263" [[package]] +name = "thiserror" +version = "2.0.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ec86235f5fcc2a73650310756d2ac5b138a5780bbbdfae3eeccec992c435ba4f" +dependencies = [ + "thiserror-impl", +] + +[[package]] +name = "thiserror-impl" +version = "2.0.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bc04cd3e1236dd4a98afca4569f2deb3f120e5422a4023be2cb683f8486292af" +dependencies = [ + "proc-macro2", + "quote", + "syn 3.0.3", +] + +[[package]] name = "tinyvec" version = "1.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" @@ -1157,6 +1206,57 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" [[package]] +name = "wasm" +version = "0.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e59fe57342dd136b22e8c7d1856e276b69e07a8f88153d218d2e54f19991a3eb" + +[[package]] +name = "wasm-bindgen" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ddb3f79143bced6de84270411622a2699cee572fc0875aeaf1e7867cf9fca1a" +dependencies = [ + "cfg-if", + "once_cell", + "rustversion", + "wasm-bindgen-macro", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-macro" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4e21a184b13fb19e157296e2c46056aec9092264fab83e4ba59e68c61b323c3d" +dependencies = [ + "quote", + "wasm-bindgen-macro-support", +] + +[[package]] +name = "wasm-bindgen-macro-support" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "fecefd9c35bd935a20fc3fc344b5f29138961e4f47fb03297d88f2587afb5ebd" +dependencies = [ + "bumpalo", + "proc-macro2", + "quote", + "syn 2.0.119", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-shared" +version = "0.2.125" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "23939e44bb9a5d7576fa2b563dc2e136628f1224e88a8deed09e04858b77871f" +dependencies = [ + "unicode-ident", +] + +[[package]] name = "windows-link" version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" diff --git a/fuzz/Cargo.toml b/fuzz/Cargo.toml @@ -41,3 +41,9 @@ name = "wallet_config" path = "fuzz_targets/wallet_config.rs" test = false doc = false + +[[bin]] +name = "vdf_proof" +path = "fuzz_targets/vdf_proof.rs" +test = false +doc = false diff --git a/fuzz/README.md b/fuzz/README.md @@ -16,6 +16,7 @@ cargo fuzz run compact_snapshot cargo fuzz run domain_json cargo fuzz run stratum_request cargo fuzz run wallet_config +cargo fuzz run vdf_proof ``` Short smoke run without installing `cargo-fuzz`: @@ -26,6 +27,7 @@ cargo run --manifest-path fuzz/Cargo.toml --bin compact_snapshot -- -runs=1 fuzz cargo run --manifest-path fuzz/Cargo.toml --bin domain_json -- -runs=1 fuzz/corpus/domain_json cargo run --manifest-path fuzz/Cargo.toml --bin stratum_request -- -runs=1 fuzz/corpus/stratum_request cargo run --manifest-path fuzz/Cargo.toml --bin wallet_config -- -runs=1 fuzz/corpus/wallet_config +cargo run --manifest-path fuzz/Cargo.toml --bin vdf_proof -- -runs=1 fuzz/corpus/vdf_proof ``` Targets intentionally accept malformed input. A parse error is fine; panics, diff --git a/fuzz/corpus/vdf_proof/chia-42-300.hex b/fuzz/corpus/vdf_proof/chia-42-300.hex @@ -0,0 +1 @@ +0000235f6d0bfcbadbd5a0d6619a8611345eb63891876d37150fdef725695ab80c6deef7684c38fe0e086355baf4786fed8a5f843d0b7a62bf1125765b016dfe965b493cfc9bcde723c5299db8db25885d130f9aef4b029f98f42831aaf53e51e33501000300d2b31e34c399ec49288e3fccb6ebaf0f3fb2e814c7c21e8579c17b5f2600b1a64d9d5b94435084b3458a9343fd1bcd3f0b9e5874556f1ab1529347b54788af1eb9268a5ee888fba85934c81b199a4228a41cb01c10b3195c95b26c17f16ff7020100 diff --git a/fuzz/fuzz_targets/vdf_proof.rs b/fuzz/fuzz_targets/vdf_proof.rs @@ -0,0 +1,47 @@ +#![no_main] + +use iuna::domain::verify_vdf; +use libfuzzer_sys::fuzz_target; + +const HEX: &[u8; 16] = b"0123456789abcdef"; + +fuzz_target!(|data: &[u8]| { + if data.is_empty() { + return; + } + + let mut proof = [0u8; 200]; + let hex_data = data.strip_suffix(b"\n").unwrap_or(data); + if hex_data.len() == proof.len() * 2 { + for (output, encoded) in proof.iter_mut().zip(hex_data.chunks_exact(2)) { + let Some(high) = decode_nibble(encoded[0]) else { + return; + }; + let Some(low) = decode_nibble(encoded[1]) else { + return; + }; + *output = high << 4 | low; + } + } else { + let copied = data.len().min(proof.len()); + proof[..copied].copy_from_slice(&data[..copied]); + } + + let mut encoded = String::with_capacity(30 + proof.len() * 2); + encoded.push_str("classgroup-wesolowski-bqfc-v1:"); + for byte in proof { + encoded.push(HEX[usize::from(byte >> 4)] as char); + encoded.push(HEX[usize::from(byte & 0x0f)] as char); + } + + let _ = verify_vdf("BBBBBBBBBBBBBBBBBBBBBBBBBBBBBBBB", 300, &encoded); +}); + +fn decode_nibble(byte: u8) -> Option<u8> { + match byte { + b'0'..=b'9' => Some(byte - b'0'), + b'a'..=b'f' => Some(byte - b'a' + 10), + b'A'..=b'F' => Some(byte - b'A' + 10), + _ => None, + } +} diff --git a/src/domain/ledger_ops.rs b/src/domain/ledger_ops.rs @@ -8,6 +8,7 @@ use super::reveal::{BurnBundleSection, canonical_burn_bundle_hashes}; use super::selection::{TransactionKind, fee_rate_key}; use super::ticket::ticket_is_eligible_for_height; use super::transaction::Transaction; +use super::vdf::vdf_solution_placeholder; use super::{ Amount, BURN_COMMITTEE_SIZE, Block, BlockSelection, BurnCommitteeMember, BurnTicket, FinalizerMode, LeaderProof, LeaderProofPayload, MINE_REWARD, OutPoint, PUBLIC_KEY_BYTES, @@ -84,7 +85,7 @@ pub(super) fn estimated_block_selection_size_bytes( finalizer_rank: 0, reward: u64::MAX, vdf_rounds: u64::MAX, - vdf_output: format!("{}:{}", "f".repeat(512), "f".repeat(512)), + vdf_output: vdf_solution_placeholder(), leader_proof: (!recovery).then(|| LeaderProof { ticket_id: "f".repeat(64), public_key: "f".repeat(64), diff --git a/src/domain/vdf.rs b/src/domain/vdf.rs @@ -1,397 +0,0 @@ -use std::{ - sync::OnceLock, - time::{Duration, Instant}, -}; - -use num_bigint::BigUint; -use num_traits::{One, Zero}; -use sha2::{Digest, Sha256}; - -use super::{Block, FinalizerMode, MAX_VDF_ROUNDS, VDF_TARGET_BLOCK_MS}; - -const VDF_RSA_2048_MODULUS_DECIMAL: &str = concat!( - "2519590847565789349402718324004839857142928212620403202777713783604366202070", - "7595556264018525880784406918290641249515082189298559149176184502808489120072", - "8449926873928072877767359714183472702618963750149718246911650776133798590957", - "0009733045974880842840179742910064245869181719511874612151517265463228221686", - "9987549182422433637259085141865462043576798423387184774447920739934236584823", - "8242811981638150106748104516603773060562016196762561338441436038339044149526", - "3443219011465754445417842402092461651572335077870774981712577246796292638635", - "6373289912154831438167899885040445364023527381951378636564391212010397122822", - "120720357", -); -const VDF_ELEMENT_HEX_LEN: usize = 512; -const VDF_CHALLENGE_MIN: u64 = 1_073_741_827; -const MIN_VDF_ROUNDS: u64 = 1; -pub(super) const VDF_RETARGET_WINDOW_BLOCKS: usize = 20; -pub(super) const MAX_VDF_RETARGET_STEP_PERCENT: u128 = 2; -pub(super) const VDF_RETARGET_DEADBAND_PERCENT: u128 = 10; -pub(super) const MIN_VDF_RETARGET_OBSERVED_BLOCK_MS: u64 = VDF_TARGET_BLOCK_MS / 4; -pub(super) const MAX_VDF_RETARGET_OBSERVED_BLOCK_MS: u64 = VDF_TARGET_BLOCK_MS * 4; - -#[derive(Clone, Copy, Debug, Eq, PartialEq)] -pub struct VdfProgress { - pub completed_steps: u64, - pub total_steps: u64, - pub completed_phase_rounds: u64, - pub phase_rounds: u64, - pub phase: VdfProgressPhase, -} - -#[derive(Clone, Copy, Debug, Eq, PartialEq)] -pub enum VdfProgressPhase { - Output, - Proof, -} - -pub fn run_vdf(seed: &str, rounds: u64) -> String { - run_vdf_with_progress(seed, rounds, Duration::MAX, |_| {}) -} - -pub fn run_vdf_with_progress( - seed: &str, - rounds: u64, - progress_interval: Duration, - mut progress: impl FnMut(VdfProgress), -) -> String { - let x = vdf_seed_element(seed); - let mut y = x.clone(); - let total_steps = rounds.saturating_mul(2); - let mut last_progress = Instant::now(); - for completed_rounds in 0..rounds { - y = square_mod(&y); - maybe_report_vdf_progress( - &mut last_progress, - progress_interval, - VdfProgress { - completed_steps: completed_rounds + 1, - total_steps, - completed_phase_rounds: completed_rounds + 1, - phase_rounds: rounds, - phase: VdfProgressPhase::Output, - }, - &mut progress, - ); - } - - let challenge = vdf_challenge_prime(seed, rounds, &y); - let proof = vdf_proof_with_progress( - &x, - rounds, - challenge, - total_steps, - &mut last_progress, - progress_interval, - &mut progress, - ); - encode_vdf_solution(y, proof) -} - -pub fn verify_vdf(seed: &str, rounds: u64, solution: &str) -> bool { - let Some((y, proof)) = decode_vdf_solution(solution) else { - return false; - }; - if y.is_zero() || y >= *vdf_modulus() || proof >= *vdf_modulus() { - return false; - } - - let x = vdf_seed_element(seed); - let challenge = vdf_challenge_prime(seed, rounds, &y); - let remainder = BigUint::from(pow_mod_small(2, rounds, challenge)); - let verified = mul_mod( - &proof.modpow(&BigUint::from(challenge), vdf_modulus()), - &x.modpow(&remainder, vdf_modulus()), - ); - verified == y -} - -pub(super) fn retarget_vdf_rounds(current_rounds: u64, observed_block_ms: u64) -> u64 { - let current = u128::from(current_rounds); - let observed = u128::from(observed_block_ms.max(1)); - let target = u128::from(VDF_TARGET_BLOCK_MS); - let deadband = target * VDF_RETARGET_DEADBAND_PERCENT / 100; - if observed >= target.saturating_sub(deadband) && observed <= target.saturating_add(deadband) { - return current_rounds; - } - - let raw_adjusted = current * target / observed; - let max_step = (current * MAX_VDF_RETARGET_STEP_PERCENT / 100).max(1); - let min_next = current - .saturating_sub(max_step) - .max(u128::from(MIN_VDF_ROUNDS)); - let max_next = current - .saturating_add(max_step) - .min(u128::from(MAX_VDF_ROUNDS)); - raw_adjusted.clamp(min_next, max_next) as u64 -} - -pub(super) fn clamped_vdf_retarget_observed_block_ms(observed_block_ms: u64) -> u64 { - observed_block_ms.clamp( - MIN_VDF_RETARGET_OBSERVED_BLOCK_MS, - MAX_VDF_RETARGET_OBSERVED_BLOCK_MS, - ) -} - -pub(super) fn vdf_retarget_observed_block_ms(parent: &Block, child: &Block) -> Option<u64> { - if child.finalizer_mode != FinalizerMode::Ticket { - return None; - } - if child.finalizer_rank != 0 { - return None; - } - - Some(clamped_vdf_retarget_observed_block_ms( - child.timestamp_ms - parent.timestamp_ms, - )) -} - -fn vdf_modulus() -> &'static BigUint { - static MODULUS: OnceLock<BigUint> = OnceLock::new(); - MODULUS.get_or_init(|| { - BigUint::parse_bytes(VDF_RSA_2048_MODULUS_DECIMAL.as_bytes(), 10) - .expect("VDF RSA-2048 modulus must parse") - }) -} - -fn vdf_seed_element(seed: &str) -> BigUint { - let one = BigUint::one(); - let two = BigUint::from(2_u32); - for attempt in 0_u32.. { - let candidate = hash_to_modulus("iuna-vdf-seed-v2", seed, attempt); - if candidate <= one { - continue; - } - let element = candidate.modpow(&two, vdf_modulus()); - if element > one { - return element; - } - } - unreachable!("VDF seed hashing must eventually produce a usable element") -} - -fn hash_to_modulus(domain: &str, seed: &str, attempt: u32) -> BigUint { - let byte_len = vdf_modulus().bits().div_ceil(8) as usize; - let mut bytes = Vec::with_capacity(byte_len); - let mut counter = 0_u32; - while bytes.len() < byte_len { - let digest = Sha256::digest(format!("{domain}:{seed}:{attempt}:{counter}").as_bytes()); - bytes.extend_from_slice(&digest); - counter = counter.saturating_add(1); - } - bytes.truncate(byte_len); - BigUint::from_bytes_be(&bytes) % vdf_modulus() -} - -fn vdf_challenge_prime(seed: &str, rounds: u64, output: &BigUint) -> u64 { - let digest = Sha256::digest(format!("iuna-vdf-challenge:{seed}:{rounds}:{output:x}")); - let mut bytes = [0_u8; 8]; - bytes.copy_from_slice(&digest[..8]); - let candidate = VDF_CHALLENGE_MIN + (u64::from_be_bytes(bytes) % VDF_CHALLENGE_MIN); - next_odd_prime(candidate | 1) -} - -fn vdf_proof_with_progress( - x: &BigUint, - rounds: u64, - challenge: u64, - total_steps: u64, - last_progress: &mut Instant, - progress_interval: Duration, - progress: &mut impl FnMut(VdfProgress), -) -> BigUint { - let mut proof = BigUint::one(); - let mut remainder = 1_u64 % challenge; - for completed_rounds in 0..rounds { - let doubled = remainder * 2; - let carry = doubled >= challenge; - proof = square_mod(&proof); - if carry { - proof = mul_mod(&proof, x); - } - remainder = doubled % challenge; - maybe_report_vdf_progress( - last_progress, - progress_interval, - VdfProgress { - completed_steps: rounds.saturating_add(completed_rounds + 1), - total_steps, - completed_phase_rounds: completed_rounds + 1, - phase_rounds: rounds, - phase: VdfProgressPhase::Proof, - }, - progress, - ); - } - proof -} - -fn maybe_report_vdf_progress( - last_progress: &mut Instant, - progress_interval: Duration, - snapshot: VdfProgress, - progress: &mut impl FnMut(VdfProgress), -) { - if snapshot.completed_steps == snapshot.total_steps - || last_progress.elapsed() >= progress_interval - { - progress(snapshot); - *last_progress = Instant::now(); - } -} - -fn encode_vdf_solution(output: BigUint, proof: BigUint) -> String { - format!( - "{output:0>width$x}:{proof:0>width$x}", - width = VDF_ELEMENT_HEX_LEN - ) -} - -fn decode_vdf_solution(solution: &str) -> Option<(BigUint, BigUint)> { - let (output, proof) = solution.split_once(':')?; - if output.len() != VDF_ELEMENT_HEX_LEN || proof.len() != VDF_ELEMENT_HEX_LEN { - return None; - } - Some(( - BigUint::parse_bytes(output.as_bytes(), 16)?, - BigUint::parse_bytes(proof.as_bytes(), 16)?, - )) -} - -fn square_mod(value: &BigUint) -> BigUint { - mul_mod(value, value) -} - -fn mul_mod(left: &BigUint, right: &BigUint) -> BigUint { - (left * right) % vdf_modulus() -} - -fn pow_mod_small(base: u64, exponent: u64, modulus: u64) -> u64 { - let mut result = 1_u128; - let mut base = u128::from(base % modulus); - let mut exponent = exponent; - let modulus = u128::from(modulus); - while exponent > 0 { - if exponent & 1 == 1 { - result = (result * base) % modulus; - } - base = (base * base) % modulus; - exponent >>= 1; - } - result as u64 -} - -fn next_odd_prime(mut candidate: u64) -> u64 { - while !is_odd_prime(candidate) { - candidate = candidate.saturating_add(2); - } - candidate -} - -fn is_odd_prime(candidate: u64) -> bool { - if candidate < 3 || candidate % 2 == 0 { - return false; - } - let mut divisor = 3_u64; - while divisor * divisor <= candidate { - if candidate % divisor == 0 { - return false; - } - divisor += 2; - } - true -} - -#[cfg(test)] -mod tests { - use std::time::Duration; - - use super::{ - VDF_ELEMENT_HEX_LEN, VdfProgressPhase, pow_mod_small, run_vdf, run_vdf_with_progress, - vdf_modulus, verify_vdf, - }; - - #[test] - fn vdf_solution_verifies_and_is_bound_to_seed_and_rounds() { - let solution = run_vdf("test-seed", 128); - - assert!(verify_vdf("test-seed", 128, &solution)); - assert!(!verify_vdf("other-seed", 128, &solution)); - assert!(!verify_vdf("test-seed", 129, &solution)); - assert!(!verify_vdf("test-seed", 128, "not-a-vdf-solution")); - } - - #[test] - fn vdf_progress_reports_output_and_proof_steps() { - let mut progress = Vec::new(); - let solution = run_vdf_with_progress("progress-seed", 4, Duration::ZERO, |snapshot| { - progress.push(snapshot); - }); - - assert!(verify_vdf("progress-seed", 4, &solution)); - assert!( - progress - .iter() - .any(|snapshot| snapshot.phase == VdfProgressPhase::Output) - ); - assert!( - progress - .iter() - .any(|snapshot| snapshot.phase == VdfProgressPhase::Proof) - ); - assert_eq!( - progress.last().map(|snapshot| snapshot.completed_steps), - Some(8) - ); - assert_eq!( - progress.last().map(|snapshot| snapshot.total_steps), - Some(8) - ); - } - - #[test] - fn vdf_solution_uses_2048_bit_elements() { - let solution = run_vdf("test-seed", 16); - let (output, proof) = solution.split_once(':').unwrap(); - - assert_eq!(output.len(), VDF_ELEMENT_HEX_LEN); - assert_eq!(proof.len(), VDF_ELEMENT_HEX_LEN); - assert!(vdf_modulus().bits() >= 2048); - } - - #[test] - fn legacy_factorable_modulus_attack_is_not_the_active_modulus() { - const LEGACY_MODULUS: u128 = 4_611_685_975_477_714_963; - const LEGACY_P: u128 = 2_147_483_629; - const LEGACY_Q: u128 = 2_147_483_647; - assert_eq!(LEGACY_P * LEGACY_Q, LEGACY_MODULUS); - assert_ne!(vdf_modulus().to_str_radix(10), LEGACY_MODULUS.to_string()); - - let phi = (LEGACY_P - 1) * (LEGACY_Q - 1); - let seed = 42_u128; - let rounds = 10_000_u64; - let sequential = legacy_repeated_squaring(seed, rounds, LEGACY_MODULUS); - let shortcut_exponent = pow_mod_small(2, rounds, phi as u64) as u128; - let shortcut = legacy_mod_pow(seed, shortcut_exponent, LEGACY_MODULUS); - - assert_eq!(shortcut, sequential); - } - - fn legacy_repeated_squaring(mut value: u128, rounds: u64, modulus: u128) -> u128 { - for _ in 0..rounds { - value = (value * value) % modulus; - } - value - } - - fn legacy_mod_pow(mut base: u128, mut exponent: u128, modulus: u128) -> u128 { - let mut result = 1_u128; - while exponent > 0 { - if exponent & 1 == 1 { - result = (result * base) % modulus; - } - base = (base * base) % modulus; - exponent >>= 1; - } - result - } -} diff --git a/src/domain/vdf/arithmetic.rs b/src/domain/vdf/arithmetic.rs @@ -0,0 +1,286 @@ +use kyn_vdf::Form; +use num_bigint::BigInt; +use num_integer::Integer; +use num_traits::{One, Signed, Zero}; + +#[cfg(test)] +pub(super) fn nudupl(form: &Form, discriminant: &BigInt, threshold: &BigInt) -> Form { + nudupl_owned(form.clone(), discriminant, threshold) +} + +pub(super) fn nudupl_owned(form: Form, discriminant: &BigInt, threshold: &BigInt) -> Form { + let two = BigInt::from(2); + let four = BigInt::from(4); + let Form { a, b, c } = form; + let mut a1 = a; + let mut c1 = c; + + let gcd = if b.is_negative() { + let b_abs = -&b; + let gcd = b_abs.extended_gcd(&a1); + (-gcd.x, gcd.gcd) + } else { + let gcd = b.extended_gcd(&a1); + (gcd.x, gcd.gcd) + }; + + let mut k = -(&gcd.0 * &c1); + let s = gcd.1; + if s != BigInt::one() { + a1 /= &s; + c1 *= &s; + } + k = mod_positive(k, &a1); + + if a1 < *threshold { + let t = &a1 * &k; + let result_a = &a1 * &a1; + let result_b = &two * &t + &b; + let result_c = ((&b + &t) * &k + &c1) / &a1; + Form::new(result_a, result_b, result_c) + } else { + let (co2, co1, _r2, r1) = xgcd_partial(&a1, &k, threshold); + let m2 = (&b * &r1 - &c1 * &co1) / &a1; + + let mut result_a = &r1 * &r1 - &co1 * &m2; + if !co1.is_negative() { + result_a = -result_a; + } + + let result_b = mod_positive( + (&two * (&a1 * &r1 - &result_a * &co2)) / &co1 - &b, + &(&result_a * &two), + ); + let mut result_c = (&result_b * &result_b - discriminant) / (&result_a * &four); + + if result_a.is_negative() { + result_a = -result_a; + result_c = -result_c; + } + + Form::new(result_a, result_b, result_c) + } +} + +pub(super) fn nucomp(left: &Form, right: &Form, discriminant: &BigInt, threshold: &BigInt) -> Form { + if left.a > right.a { + return nucomp(right, left, discriminant, threshold); + } + + let two = BigInt::from(2); + let four = BigInt::from(4); + let mut a1 = left.a.clone(); + let mut a2 = right.a.clone(); + let mut c2 = right.c.clone(); + let ss = (&left.b + &right.b) / &two; + let m = (&left.b - &right.b) / &two; + + let t = &a2 % &a1; + let (v1, sp) = if t.is_zero() { + (BigInt::zero(), a1.clone()) + } else { + let gcd = t.extended_gcd(&a1); + (gcd.x, gcd.gcd) + }; + let mut k = mod_positive(&m * &v1, &a1); + + if sp != BigInt::one() { + let gcd = ss.extended_gcd(&sp); + let v2 = gcd.x; + let u2 = gcd.y; + let s = gcd.gcd; + k = &k * &u2 - &v2 * &c2; + if s != BigInt::one() { + a1 /= &s; + a2 /= &s; + c2 *= &s; + } + k = mod_positive(k, &a1); + } + + if a1 < *threshold { + let t = &a2 * &k; + let result_a = &a2 * &a1; + let result_b = &two * &t + &right.b; + let result_c = ((&right.b + &t) * &k + &c2) / &a1; + Form::new(result_a, result_b, result_c) + } else { + let (co2, co1, _r2, r1) = xgcd_partial(&a1, &k, threshold); + let m1 = (&m * &co1 + &a2 * &r1) / &a1; + let m2 = (&ss * &r1 - &c2 * &co1) / &a1; + + let mut result_a = &r1 * &m1 - &co1 * &m2; + if !co1.is_negative() { + result_a = -result_a; + } + + let t = &a2 * &r1; + let result_b = mod_positive( + (&two * (&t - &result_a * &co2)) / &co1 - &right.b, + &(&result_a * &two), + ); + let mut result_c = (&result_b * &result_b - discriminant) / (&result_a * &four); + + if result_a.is_negative() { + result_a = -result_a; + result_c = -result_c; + } + + Form::new(result_a, result_b, result_c) + } +} + +pub(super) fn xgcd_partial( + r2: &BigInt, + r1: &BigInt, + threshold: &BigInt, +) -> (BigInt, BigInt, BigInt, BigInt) { + let mut r2 = r2.clone(); + let mut r1 = r1.clone(); + let mut co2 = BigInt::zero(); + let mut co1 = BigInt::from(-1); + + while !r1.is_zero() && &r1 > threshold { + let bits = r2.bits().max(r1.bits()).saturating_sub(63); + let mut rr2 = shifted_low_word(&r2, bits); + let mut rr1 = shifted_low_word(&r1, bits); + let threshold_word = shifted_low_word(threshold, bits); + + let mut aa2 = 0_i128; + let mut aa1 = 1_i128; + let mut bb2 = 1_i128; + let mut bb1 = 0_i128; + let mut steps = 0_u32; + + while rr1 != 0 && rr1 > threshold_word { + let q = rr2 / rr1; + let next_r = rr2 - q * rr1; + let next_a = aa2 - q * aa1; + let next_b = bb2 - q * bb1; + + if steps & 1 == 1 { + if next_r < -next_b || rr1 - next_r < next_a - aa1 { + break; + } + } else if next_r < -next_a || rr1 - next_r < next_b - bb1 { + break; + } + + rr2 = rr1; + rr1 = next_r; + aa2 = aa1; + aa1 = next_a; + bb2 = bb1; + bb1 = next_b; + steps += 1; + } + + if steps == 0 { + let (q, next_r) = r2.div_rem(&r1); + let next_co = &co2 - &q * &co1; + r2 = r1; + r1 = next_r; + co2 = co1; + co1 = next_co; + } else { + let old_r2 = r2; + let old_r1 = r1; + r2 = scaled(&old_r2, bb2) + scaled(&old_r1, aa2); + r1 = scaled(&old_r1, aa1) + scaled(&old_r2, bb1); + + let old_co2 = co2; + let old_co1 = co1; + co2 = scaled(&old_co2, bb2) + scaled(&old_co1, aa2); + co1 = scaled(&old_co1, aa1) + scaled(&old_co2, bb1); + + if r1.is_negative() { + r1 = -r1; + co1 = -co1; + } + if r2.is_negative() { + r2 = -r2; + co2 = -co2; + } + } + } + + if r2.is_negative() { + r2 = -r2; + co2 = -co2; + co1 = -co1; + } + + (co2, co1, r2, r1) +} + +fn scaled(value: &BigInt, scalar: i128) -> BigInt { + value * BigInt::from(scalar) +} + +fn mod_positive(mut value: BigInt, modulus: &BigInt) -> BigInt { + value %= modulus; + if value.is_negative() { + value += modulus; + } + value +} + +fn shifted_low_word(value: &BigInt, shift_bits: u64) -> i128 { + debug_assert!(!value.is_negative()); + let mut digits = value.iter_u64_digits(); + let digit_index = usize::try_from(shift_bits / 64).unwrap_or(usize::MAX); + let offset = (shift_bits % 64) as u32; + let low = digits.nth(digit_index).unwrap_or(0); + if offset == 0 { + return i128::from(low); + } + let high = digits.next().unwrap_or(0); + i128::from((low >> offset) | (high << (64 - offset))) +} + +#[cfg(test)] +mod tests { + use kyn_vdf::{Form, create_discriminant, isqrt_fourth}; + use num_traits::Signed; + + use super::{nucomp, nudupl}; + + #[test] + fn optimized_nudupl_matches_kyn_across_sequential_squares() { + let discriminant = create_discriminant(b"iuna-vdf-arithmetic-nudupl", 1024).unwrap(); + let threshold = isqrt_fourth(&discriminant.abs()); + let mut form = Form::generator(&discriminant).unwrap(); + + for round in 1..=10_000 { + let mut expected = form.nudupl(&discriminant, &threshold); + expected.reduce(&discriminant); + let mut actual = nudupl(&form, &discriminant, &threshold); + actual.reduce(&discriminant); + + assert_eq!(actual, expected, "round={round}"); + form = actual; + } + } + + #[test] + fn optimized_nucomp_matches_kyn_across_sequential_compositions() { + let discriminant = create_discriminant(b"iuna-vdf-arithmetic-nucomp", 1024).unwrap(); + let threshold = isqrt_fourth(&discriminant.abs()); + let generator = Form::generator(&discriminant).unwrap(); + let mut left = generator.clone(); + let mut right = generator.nudupl(&discriminant, &threshold); + right.reduce(&discriminant); + + for round in 1..=2_000 { + let mut expected = left.nucomp(&right, &discriminant, &threshold); + expected.reduce(&discriminant); + let mut actual = nucomp(&left, &right, &discriminant, &threshold); + actual.reduce(&discriminant); + + assert_eq!(actual, expected, "round={round}"); + left = actual; + right = right.nudupl(&discriminant, &threshold); + right.reduce(&discriminant); + } + } +} diff --git a/src/domain/vdf/limb_arithmetic.rs b/src/domain/vdf/limb_arithmetic.rs @@ -0,0 +1,570 @@ +use std::cmp::Ordering; + +use kyn_vdf::Form; +use num_bigint::BigInt; + +use super::limbs::{LimbInt, LimbScratch, xgcd_partial_with_scratch}; + +#[derive(Clone, Debug, Eq, PartialEq)] +pub(super) struct LimbForm { + a: LimbInt, + b: LimbInt, + c: LimbInt, +} + +#[derive(Default)] +pub(super) struct LimbFormScratch { + limbs: LimbScratch, + k_product: LimbInt, + work0: LimbInt, + work1: LimbInt, + modulus: LimbInt, +} + +impl LimbForm { + pub(super) fn identity(discriminant: &LimbInt) -> Self { + let one = LimbInt::one(); + let four = LimbInt::from_u64(4); + Self { + a: one.clone(), + b: one.clone(), + c: one.sub(discriminant).div(&four), + } + } + + pub(super) fn from_form(form: &Form) -> Self { + Self { + a: to_limb(&form.a), + b: to_limb(&form.b), + c: to_limb(&form.c), + } + } + + pub(super) fn into_form(self) -> Form { + Form::new(from_limb(&self.a), from_limb(&self.b), from_limb(&self.c)) + } + + #[cfg(test)] + pub(super) fn nudupl_reduce(self, discriminant: &LimbInt, threshold: &LimbInt) -> Self { + let mut scratch = LimbFormScratch::default(); + self.nudupl_reduce_with_scratch(discriminant, threshold, &mut scratch) + } + + pub(super) fn nudupl_reduce_with_scratch( + self, + discriminant: &LimbInt, + threshold: &LimbInt, + scratch: &mut LimbFormScratch, + ) -> Self { + if self.is_identity() { + return self; + } + let mut form = nudupl_owned(self, discriminant, threshold, scratch); + reduce_with_scratch(&mut form, scratch); + form + } + + #[cfg(test)] + pub(super) fn nucomp_reduce( + &self, + other: &Self, + discriminant: &LimbInt, + threshold: &LimbInt, + ) -> Self { + let mut scratch = LimbFormScratch::default(); + self.nucomp_reduce_with_scratch(other, discriminant, threshold, &mut scratch) + } + + pub(super) fn nucomp_reduce_with_scratch( + &self, + other: &Self, + discriminant: &LimbInt, + threshold: &LimbInt, + scratch: &mut LimbFormScratch, + ) -> Self { + if self.is_identity() { + return other.clone(); + } + if other.is_identity() { + return self.clone(); + } + let mut form = nucomp(self, other, discriminant, threshold, scratch); + reduce_with_scratch(&mut form, scratch); + form + } + + pub(super) fn compose_unreduced( + &self, + other: &Self, + discriminant: &LimbInt, + threshold: &LimbInt, + scratch: &mut LimbFormScratch, + ) -> Self { + if self.is_identity() { + return other.clone(); + } + if other.is_identity() { + return self.clone(); + } + nucomp(self, other, discriminant, threshold, scratch) + } + + pub(super) fn fast_pow_u64_with_scratch( + &self, + exponent: u64, + discriminant: &LimbInt, + threshold: &LimbInt, + scratch: &mut LimbFormScratch, + ) -> Self { + if exponent == 0 { + return Self::identity(discriminant); + } + if self.is_identity() { + return self.clone(); + } + + let mut result = self.clone(); + let max_bits = discriminant.bit_len() / 2; + let num_bits = u64::BITS - exponent.leading_zeros(); + + for bit in (0..num_bits.saturating_sub(1)).rev() { + result = nudupl_owned(result, discriminant, threshold, scratch); + if result.a.bit_len() > max_bits { + reduce_with_scratch(&mut result, scratch); + } + + if ((exponent >> bit) & 1) == 1 { + result = nucomp(&result, self, discriminant, threshold, scratch); + } + } + + reduce_with_scratch(&mut result, scratch); + result + } + + pub(super) fn reduce(&mut self) { + let mut scratch = LimbFormScratch::default(); + reduce_with_scratch(self, &mut scratch); + } + + fn is_identity(&self) -> bool { + self.a.is_one() && self.b.is_one() + } +} + +pub(super) fn to_limb(value: &BigInt) -> LimbInt { + LimbInt::from_bigint(value) +} + +pub(super) fn from_limb(value: &LimbInt) -> BigInt { + value.to_bigint() +} + +fn nudupl_owned( + form: LimbForm, + discriminant: &LimbInt, + threshold: &LimbInt, + scratch: &mut LimbFormScratch, +) -> LimbForm { + let LimbForm { a, b, c } = form; + let mut a1 = a; + let mut c1 = c; + + let gcd = if b.is_negative() { + let b_abs = b.clone().negated(); + let gcd = b_abs.left_extended_gcd_positive_with_scratch(&a1, &mut scratch.limbs); + (gcd.x.negated(), gcd.gcd) + } else { + let gcd = b.left_extended_gcd_positive_with_scratch(&a1, &mut scratch.limbs); + (gcd.x, gcd.gcd) + }; + + gcd.0.mul_into(&c1, &mut scratch.k_product); + scratch.k_product.negate_assign(); + let s = gcd.1; + if !s.is_one() { + a1 = a1.div_with_scratch(&s, &mut scratch.limbs); + c1 = c1.mul(&s); + } + let k = scratch + .k_product + .mod_positive_with_scratch(&a1, &mut scratch.limbs); + + if a1.cmp(threshold) == Ordering::Less { + let t = a1.mul(&k); + let result_a = a1.square(); + let result_b = t.shl_bits(1).add_owned(&b); + let result_c = b + .add(&t) + .mul(&k) + .add_owned(&c1) + .div_with_scratch(&a1, &mut scratch.limbs); + LimbForm { + a: result_a, + b: result_b, + c: result_c, + } + } else { + let (co2, co1, _r2, r1) = xgcd_partial_with_scratch(&a1, &k, threshold, &mut scratch.limbs); + b.mul_into(&r1, &mut scratch.work0); + c1.mul_into(&co1, &mut scratch.work1); + scratch.work0.sub_assign(&scratch.work1); + let m2 = scratch.work0.div_with_scratch(&a1, &mut scratch.limbs); + + r1.square_into(&mut scratch.work0); + co1.mul_into(&m2, &mut scratch.work1); + scratch.work0.sub_assign(&scratch.work1); + let mut result_a = std::mem::take(&mut scratch.work0); + if !co1.is_negative() { + result_a = result_a.negated(); + } + + a1.mul_into(&r1, &mut scratch.work0); + result_a.mul_into(&co2, &mut scratch.work1); + scratch.work0.sub_assign(&scratch.work1); + scratch.work0.shl_bits_assign(1); + let mut result_b = scratch.work0.div_with_scratch(&co1, &mut scratch.limbs); + result_b.sub_assign(&b); + result_a.shl_bits_into(1, &mut scratch.modulus); + let result_b = result_b.mod_positive_with_scratch(&scratch.modulus, &mut scratch.limbs); + + result_b.square_into(&mut scratch.work0); + scratch.work0.sub_assign(discriminant); + result_a.shl_bits_into(2, &mut scratch.modulus); + let mut result_c = scratch + .work0 + .div_with_scratch(&scratch.modulus, &mut scratch.limbs); + + if result_a.is_negative() { + result_a = result_a.negated(); + result_c = result_c.negated(); + } + + LimbForm { + a: result_a, + b: result_b, + c: result_c, + } + } +} + +fn nucomp( + left: &LimbForm, + right: &LimbForm, + discriminant: &LimbInt, + threshold: &LimbInt, + scratch: &mut LimbFormScratch, +) -> LimbForm { + if left.a.cmp(&right.a) == Ordering::Greater { + return nucomp(right, left, discriminant, threshold, scratch); + } + + let mut a1 = left.a.clone(); + let mut a2 = right.a.clone(); + let mut c2 = right.c.clone(); + let ss = left.b.add(&right.b).div2_exact(); + let m = left.b.sub(&right.b).div2_exact(); + + let t = a2.rem_with_scratch(&a1, &mut scratch.limbs); + let (v1, sp) = if t.is_zero() { + (LimbInt::zero(), a1.clone()) + } else { + let gcd = t.left_extended_gcd_positive_with_scratch(&a1, &mut scratch.limbs); + (gcd.x, gcd.gcd) + }; + let mut k = m + .mul(&v1) + .mod_positive_with_scratch(&a1, &mut scratch.limbs); + + if !sp.is_one() { + let gcd = ss.extended_gcd_with_scratch(&sp, &mut scratch.limbs); + let v2 = gcd.x; + let u2 = gcd.y; + let s = gcd.gcd; + k = k.mul(&u2).sub_owned(&v2.mul(&c2)); + if !s.is_one() { + a1 = a1.div_with_scratch(&s, &mut scratch.limbs); + a2 = a2.div_with_scratch(&s, &mut scratch.limbs); + c2 = c2.mul(&s); + } + k = k.mod_positive_with_scratch(&a1, &mut scratch.limbs); + } + + if a1.cmp(threshold) == Ordering::Less { + let t = a2.mul(&k); + let result_a = a2.mul(&a1); + let result_b = t.shl_bits(1).add_owned(&right.b); + let result_c = right + .b + .add(&t) + .mul(&k) + .add_owned(&c2) + .div_with_scratch(&a1, &mut scratch.limbs); + LimbForm { + a: result_a, + b: result_b, + c: result_c, + } + } else { + let (co2, co1, _r2, r1) = xgcd_partial_with_scratch(&a1, &k, threshold, &mut scratch.limbs); + let m1 = m + .mul(&co1) + .add_owned(&a2.mul(&r1)) + .div_with_scratch(&a1, &mut scratch.limbs); + let m2 = ss + .mul(&r1) + .sub_owned(&c2.mul(&co1)) + .div_with_scratch(&a1, &mut scratch.limbs); + + let mut result_a = r1.mul(&m1).sub_owned(&co1.mul(&m2)); + if !co1.is_negative() { + result_a = result_a.negated(); + } + + let t = a2.mul(&r1); + let result_b = t + .sub_owned(&result_a.mul(&co2)) + .shl_bits_owned(1) + .div_with_scratch(&co1, &mut scratch.limbs) + .sub_owned(&right.b) + .mod_positive_with_scratch(&result_a.shl_bits(1), &mut scratch.limbs); + let mut result_c = result_b + .square() + .sub_owned(discriminant) + .div_with_scratch(&result_a.shl_bits(2), &mut scratch.limbs); + + if result_a.is_negative() { + result_a = result_a.negated(); + result_c = result_c.negated(); + } + + LimbForm { + a: result_a, + b: result_b, + c: result_c, + } + } +} + +#[inline(always)] +fn reduce_with_scratch(form: &mut LimbForm, scratch: &mut LimbFormScratch) { + while !finish_if_reduced(form) { + reduce_once(form, scratch); + } +} + +#[inline(always)] +fn finish_if_reduced(form: &mut LimbForm) -> bool { + if form.a.abs_cmp(&form.b) == Ordering::Less || form.c.abs_cmp(&form.b) == Ordering::Less { + return false; + } + + match form.a.cmp(&form.c) { + Ordering::Greater => { + std::mem::swap(&mut form.a, &mut form.c); + form.b = std::mem::take(&mut form.b).negated(); + } + Ordering::Equal if form.b.is_negative() => { + form.b = std::mem::take(&mut form.b).negated(); + } + _ => {} + } + true +} + +#[inline(always)] +fn reduce_once(form: &mut LimbForm, scratch: &mut LimbFormScratch) { + form.c.shl_bits_into(1, &mut scratch.modulus); + form.b.add_into(&form.c, &mut scratch.work0); + + if scratch.work0.is_zero() + || (!scratch.work0.is_negative() + && scratch.work0.abs_cmp(&scratch.modulus) == Ordering::Less) + { + let old_a = std::mem::take(&mut form.a); + let old_b = std::mem::take(&mut form.b); + let old_c = std::mem::take(&mut form.c); + form.a = old_c; + form.b = old_b.negated(); + form.c = old_a; + return; + } + + if scratch.work0.is_negative() { + if scratch.work0.abs_cmp(&scratch.modulus) != Ordering::Greater { + let old_a = std::mem::take(&mut form.a); + let old_b = std::mem::take(&mut form.b); + let old_c = std::mem::take(&mut form.c); + + old_c.shl_bits_into(1, &mut scratch.work0); + scratch.work0.add_assign(&old_b); + old_c.add_into(&old_b, &mut scratch.work1); + + form.a = old_c; + form.b = std::mem::take(&mut scratch.work0).negated(); + form.c = old_a; + form.c.add_assign(&scratch.work1); + return; + } + } else if scratch.work0.abs_cmp_double(&scratch.modulus) == Ordering::Less { + let old_a = std::mem::take(&mut form.a); + let old_b = std::mem::take(&mut form.b); + let old_c = std::mem::take(&mut form.c); + + old_c.shl_bits_into(1, &mut scratch.work0); + scratch.work0.sub_assign(&old_b); + old_c.sub_into(&old_b, &mut scratch.work1); + + form.a = old_c; + form.b = std::mem::take(&mut scratch.work0); + form.c = old_a; + form.c.add_assign(&scratch.work1); + return; + } + + let s = scratch + .work0 + .div_floor_with_scratch(&scratch.modulus, &mut scratch.limbs); + let old_a = std::mem::take(&mut form.a); + let old_b = std::mem::take(&mut form.b); + let old_c = std::mem::take(&mut form.c); + + form.a = old_c; + form.a.mul_into(&s, &mut scratch.work0); + scratch.work0.sub_into(&old_b, &mut scratch.work1); + + form.b = std::mem::take(&mut scratch.work0); + form.b.add_assign(&scratch.work1); + + s.mul_into(&scratch.work1, &mut scratch.work0); + form.c = old_a; + form.c.add_assign(&scratch.work0); +} + +#[cfg(test)] +mod tests { + use std::time::Instant; + + use kyn_vdf::{Form, create_discriminant, isqrt_fourth}; + use num_traits::Signed; + + use super::{LimbForm, LimbFormScratch, to_limb}; + + #[test] + fn limb_nudupl_matches_kyn_across_sequential_squares() { + let discriminant = create_discriminant(b"iuna-vdf-limb-nudupl", 1024).unwrap(); + let threshold = isqrt_fourth(&discriminant.abs()); + let limb_discriminant = to_limb(&discriminant); + let limb_threshold = to_limb(&threshold); + let mut expected = Form::generator(&discriminant).unwrap(); + let mut actual = LimbForm::from_form(&expected); + + for round in 1..=10_000 { + expected = expected.nudupl(&discriminant, &threshold); + expected.reduce(&discriminant); + actual = actual.nudupl_reduce(&limb_discriminant, &limb_threshold); + + assert_eq!(actual.clone().into_form(), expected, "round={round}"); + } + } + + #[test] + fn limb_nucomp_matches_kyn_across_sequential_compositions() { + let discriminant = create_discriminant(b"iuna-vdf-limb-nucomp", 1024).unwrap(); + let threshold = isqrt_fourth(&discriminant.abs()); + let limb_discriminant = to_limb(&discriminant); + let limb_threshold = to_limb(&threshold); + let generator = Form::generator(&discriminant).unwrap(); + let mut expected_left = generator.clone(); + let mut actual_left = LimbForm::from_form(&expected_left); + let mut expected_right = generator.nudupl(&discriminant, &threshold); + expected_right.reduce(&discriminant); + let mut actual_right = LimbForm::from_form(&expected_right); + + for round in 1..=2_000 { + let mut expected = expected_left.nucomp(&expected_right, &discriminant, &threshold); + expected.reduce(&discriminant); + let actual = + actual_left.nucomp_reduce(&actual_right, &limb_discriminant, &limb_threshold); + + assert_eq!(actual.clone().into_form(), expected, "round={round}"); + expected_left = expected; + actual_left = actual; + expected_right = expected_right.nudupl(&discriminant, &threshold); + expected_right.reduce(&discriminant); + actual_right = actual_right.nudupl_reduce(&limb_discriminant, &limb_threshold); + } + } + + #[test] + #[ignore = "manual custom limb NUDUPL benchmark"] + fn benchmark_limb_nudupl_against_optimized_bigint() { + let rounds = 100_000; + let discriminant = create_discriminant(b"iuna-vdf-limb-benchmark", 1024).unwrap(); + let threshold = isqrt_fourth(&discriminant.abs()); + let limb_discriminant = to_limb(&discriminant); + let limb_threshold = to_limb(&threshold); + let generator = Form::generator(&discriminant).unwrap(); + + let mut bigint_output = generator.clone(); + let started = Instant::now(); + for _ in 0..rounds { + bigint_output = crate::domain::vdf::arithmetic::nudupl_owned( + bigint_output, + &discriminant, + &threshold, + ); + crate::domain::vdf::reducer::reduce(&mut bigint_output); + } + let bigint_elapsed = started.elapsed(); + + let mut limb_output = LimbForm::from_form(&generator); + let mut limb_scratch = LimbFormScratch::default(); + let started = Instant::now(); + for _ in 0..rounds { + limb_output = limb_output.nudupl_reduce_with_scratch( + &limb_discriminant, + &limb_threshold, + &mut limb_scratch, + ); + } + let limb_elapsed = started.elapsed(); + + assert_eq!(limb_output.into_form(), bigint_output); + eprintln!( + "rounds={rounds} optimized_bigint={bigint_elapsed:?} custom_limb={limb_elapsed:?} speedup={:.2}x", + bigint_elapsed.as_secs_f64() / limb_elapsed.as_secs_f64() + ); + } + + #[test] + #[ignore = "manual custom limb NUDUPL phase benchmark"] + fn benchmark_limb_nudupl_phases() { + let rounds = 100_000; + let discriminant = create_discriminant(b"iuna-vdf-limb-benchmark", 1024).unwrap(); + let threshold = isqrt_fourth(&discriminant.abs()); + let limb_discriminant = to_limb(&discriminant); + let limb_threshold = to_limb(&threshold); + let generator = Form::generator(&discriminant).unwrap(); + let mut output = LimbForm::from_form(&generator); + let mut scratch = LimbFormScratch::default(); + let mut nudupl_elapsed = std::time::Duration::ZERO; + let mut reduce_elapsed = std::time::Duration::ZERO; + + for _ in 0..rounds { + let started = Instant::now(); + output = super::nudupl_owned(output, &limb_discriminant, &limb_threshold, &mut scratch); + nudupl_elapsed += started.elapsed(); + + let started = Instant::now(); + super::reduce_with_scratch(&mut output, &mut scratch); + reduce_elapsed += started.elapsed(); + } + + assert!(output.clone().into_form().is_reduced()); + eprintln!( + "rounds={rounds} nudupl={nudupl_elapsed:?} reduce={reduce_elapsed:?} total={:?}", + nudupl_elapsed + reduce_elapsed + ); + } +} diff --git a/src/domain/vdf/limbs/mod.rs b/src/domain/vdf/limbs/mod.rs @@ -0,0 +1,2413 @@ +#![allow(dead_code)] + +use std::cmp::Ordering; + +use num_bigint::{BigInt, Sign as BigSign}; +use num_traits::Signed; + +type LimbVec = Vec<u64>; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum Sign { + Negative, + Zero, + Positive, +} + +#[derive(Clone, Debug, Eq, PartialEq)] +pub(super) struct LimbInt { + sign: Sign, + limbs: LimbVec, +} + +impl Default for LimbInt { + fn default() -> Self { + Self::zero() + } +} + +pub(super) struct ExtendedGcd { + pub(super) x: LimbInt, + pub(super) y: LimbInt, + pub(super) gcd: LimbInt, +} + +pub(super) struct LeftExtendedGcd { + pub(super) x: LimbInt, + pub(super) gcd: LimbInt, +} + +#[derive(Default)] +pub(super) struct LimbScratch { + division: DivisionScratch, + linear_left: LimbVec, + linear_right: LimbVec, +} + +#[derive(Default)] +struct DivisionScratch { + normalized_numerator: LimbVec, + normalized_denominator: LimbVec, + quotient: LimbVec, + remainder: LimbVec, +} + +impl LimbInt { + pub(super) fn zero() -> Self { + Self { + sign: Sign::Zero, + limbs: LimbVec::new(), + } + } + + pub(super) fn one() -> Self { + Self { + sign: Sign::Positive, + limbs: vec![1], + } + } + + pub(super) fn from_u64(value: u64) -> Self { + if value == 0 { + Self::zero() + } else { + Self { + sign: Sign::Positive, + limbs: vec![value], + } + } + } + + pub(super) fn from_i128(value: i128) -> Self { + if value == 0 { + return Self::zero(); + } + + let sign = if value < 0 { + Sign::Negative + } else { + Sign::Positive + }; + let mut magnitude = value.unsigned_abs(); + let mut limbs = LimbVec::new(); + while magnitude > 0 { + limbs.push(magnitude as u64); + magnitude >>= 64; + } + Self { sign, limbs } + } + + pub(super) fn from_bigint(value: &BigInt) -> Self { + if value == &BigInt::from(0) { + return Self::zero(); + } + + let sign = if value.is_negative() { + Sign::Negative + } else { + Sign::Positive + }; + let magnitude = value + .abs() + .to_biguint() + .expect("absolute BigInt is non-negative"); + let mut limbs: LimbVec = magnitude.iter_u64_digits().collect(); + trim_leading_zero_limbs(&mut limbs); + Self { sign, limbs } + } + + pub(super) fn to_bigint(&self) -> BigInt { + let sign = match self.sign { + Sign::Negative => BigSign::Minus, + Sign::Zero => BigSign::NoSign, + Sign::Positive => BigSign::Plus, + }; + BigInt::from_biguint(sign, self.abs_biguint()) + } + + #[inline(always)] + pub(super) fn is_zero(&self) -> bool { + self.sign == Sign::Zero + } + + #[inline(always)] + pub(super) fn is_negative(&self) -> bool { + self.sign == Sign::Negative + } + + #[inline(always)] + pub(super) fn is_one(&self) -> bool { + self.sign == Sign::Positive && self.limbs.as_slice() == [1] + } + + #[inline(always)] + pub(super) fn is_minus_one(&self) -> bool { + self.sign == Sign::Negative && self.limbs.as_slice() == [1] + } + + #[inline(always)] + pub(super) fn bit_len(&self) -> u64 { + let Some(last) = self.limbs.last() else { + return 0; + }; + let top_bits = u64::BITS - last.leading_zeros(); + ((self.limbs.len() as u64 - 1) * 64) + u64::from(top_bits) + } + + #[inline(always)] + pub(super) fn abs_cmp(&self, other: &Self) -> Ordering { + cmp_abs_limbs(&self.limbs, &other.limbs) + } + + #[inline(always)] + pub(super) fn abs_cmp_double(&self, other: &Self) -> Ordering { + cmp_abs_to_double(&self.limbs, &other.limbs) + } + + pub(super) fn cmp(&self, other: &Self) -> Ordering { + match (self.sign, other.sign) { + (Sign::Negative, Sign::Negative) => cmp_abs_limbs(&other.limbs, &self.limbs), + (Sign::Negative, _) => Ordering::Less, + (_, Sign::Negative) => Ordering::Greater, + (Sign::Zero, Sign::Zero) => Ordering::Equal, + (Sign::Zero, Sign::Positive) => Ordering::Less, + (Sign::Positive, Sign::Zero) => Ordering::Greater, + (Sign::Positive, Sign::Positive) => cmp_abs_limbs(&self.limbs, &other.limbs), + } + } + + pub(super) fn negated(mut self) -> Self { + self.sign = match self.sign { + Sign::Negative => Sign::Positive, + Sign::Zero => Sign::Zero, + Sign::Positive => Sign::Negative, + }; + self + } + + pub(super) fn negate_assign(&mut self) { + self.sign = match self.sign { + Sign::Negative => Sign::Positive, + Sign::Zero => Sign::Zero, + Sign::Positive => Sign::Negative, + }; + } + + pub(super) fn add(&self, other: &Self) -> Self { + match (self.sign, other.sign) { + (Sign::Zero, _) => other.clone(), + (_, Sign::Zero) => self.clone(), + (Sign::Positive, Sign::Positive) => { + Self::from_parts(Sign::Positive, add_abs_limbs(&self.limbs, &other.limbs)) + } + (Sign::Negative, Sign::Negative) => { + Self::from_parts(Sign::Negative, add_abs_limbs(&self.limbs, &other.limbs)) + } + (Sign::Positive, Sign::Negative) => { + subtract_signed_abs(&self.limbs, Sign::Positive, &other.limbs, Sign::Negative) + } + (Sign::Negative, Sign::Positive) => { + subtract_signed_abs(&other.limbs, Sign::Positive, &self.limbs, Sign::Negative) + } + } + } + + pub(super) fn add_owned(mut self, other: &Self) -> Self { + self.add_assign(other); + self + } + + #[inline(always)] + pub(super) fn add_into(&self, other: &Self, output: &mut Self) { + match (self.sign, other.sign) { + (Sign::Zero, _) => { + *output = other.clone(); + } + (_, Sign::Zero) => { + *output = self.clone(); + } + (Sign::Positive, Sign::Positive) | (Sign::Negative, Sign::Negative) => { + add_abs_limbs_into(&self.limbs, &other.limbs, &mut output.limbs); + output.sign = if output.limbs.is_empty() { + Sign::Zero + } else { + self.sign + }; + } + (Sign::Positive, Sign::Negative) | (Sign::Negative, Sign::Positive) => { + let sign = combine_signed_abs_limbs_into( + self.sign, + &self.limbs, + other.sign, + &other.limbs, + &mut output.limbs, + ); + output.sign = if output.limbs.is_empty() { + Sign::Zero + } else { + sign + }; + } + } + } + + pub(super) fn add_assign(&mut self, other: &Self) { + match (self.sign, other.sign) { + (Sign::Zero, _) => { + *self = other.clone(); + } + (_, Sign::Zero) => {} + (Sign::Positive, Sign::Positive) | (Sign::Negative, Sign::Negative) => { + add_abs_limbs_assign(&mut self.limbs, &other.limbs); + } + (Sign::Positive, Sign::Negative) => match cmp_abs_limbs(&self.limbs, &other.limbs) { + Ordering::Greater => sub_abs_limbs_assign(&mut self.limbs, &other.limbs), + Ordering::Less => { + self.limbs = sub_abs_limbs(&other.limbs, &self.limbs); + self.sign = Sign::Negative; + } + Ordering::Equal => *self = Self::zero(), + }, + (Sign::Negative, Sign::Positive) => match cmp_abs_limbs(&self.limbs, &other.limbs) { + Ordering::Greater => sub_abs_limbs_assign(&mut self.limbs, &other.limbs), + Ordering::Less => { + self.limbs = sub_abs_limbs(&other.limbs, &self.limbs); + self.sign = Sign::Positive; + } + Ordering::Equal => *self = Self::zero(), + }, + } + } + + pub(super) fn sub(&self, other: &Self) -> Self { + let mut result = self.clone(); + result.sub_assign(other); + result + } + + pub(super) fn sub_owned(mut self, other: &Self) -> Self { + self.sub_assign(other); + self + } + + #[inline(always)] + pub(super) fn sub_into(&self, other: &Self, output: &mut Self) { + match (self.sign, other.sign) { + (_, Sign::Zero) => { + *output = self.clone(); + } + (Sign::Zero, _) => { + *output = other.clone().negated(); + } + (Sign::Positive, Sign::Positive) | (Sign::Negative, Sign::Negative) => { + match cmp_abs_limbs(&self.limbs, &other.limbs) { + Ordering::Greater => { + sub_abs_limbs_into(&self.limbs, &other.limbs, &mut output.limbs); + output.sign = self.sign; + } + Ordering::Less => { + sub_abs_limbs_into(&other.limbs, &self.limbs, &mut output.limbs); + output.sign = match self.sign { + Sign::Positive => Sign::Negative, + Sign::Negative => Sign::Positive, + Sign::Zero => Sign::Zero, + }; + } + Ordering::Equal => { + *output = Self::zero(); + } + } + } + (Sign::Positive, Sign::Negative) | (Sign::Negative, Sign::Positive) => { + add_abs_limbs_into(&self.limbs, &other.limbs, &mut output.limbs); + output.sign = if output.limbs.is_empty() { + Sign::Zero + } else { + self.sign + }; + } + } + } + + pub(super) fn sub_assign(&mut self, other: &Self) { + match (self.sign, other.sign) { + (_, Sign::Zero) => {} + (Sign::Zero, _) => { + *self = other.clone().negated(); + } + (Sign::Positive, Sign::Positive) => match cmp_abs_limbs(&self.limbs, &other.limbs) { + Ordering::Greater => sub_abs_limbs_assign(&mut self.limbs, &other.limbs), + Ordering::Less => { + self.limbs = sub_abs_limbs(&other.limbs, &self.limbs); + self.sign = Sign::Negative; + } + Ordering::Equal => *self = Self::zero(), + }, + (Sign::Negative, Sign::Negative) => match cmp_abs_limbs(&self.limbs, &other.limbs) { + Ordering::Greater => sub_abs_limbs_assign(&mut self.limbs, &other.limbs), + Ordering::Less => { + self.limbs = sub_abs_limbs(&other.limbs, &self.limbs); + self.sign = Sign::Positive; + } + Ordering::Equal => *self = Self::zero(), + }, + (Sign::Positive, Sign::Negative) | (Sign::Negative, Sign::Positive) => { + add_abs_limbs_assign(&mut self.limbs, &other.limbs); + } + } + } + + pub(super) fn mul(&self, other: &Self) -> Self { + if self.is_zero() || other.is_zero() { + return Self::zero(); + } + + let sign = if self.sign == other.sign { + Sign::Positive + } else { + Sign::Negative + }; + let limbs = match (self.limbs.as_slice(), other.limbs.as_slice()) { + ([scalar], value) => mul_abs_one_limb(value, *scalar), + (value, [scalar]) => mul_abs_one_limb(value, *scalar), + _ => mul_abs_limbs(&self.limbs, &other.limbs), + }; + Self::from_parts(sign, limbs) + } + + pub(super) fn mul_into(&self, other: &Self, output: &mut Self) { + if self.is_zero() || other.is_zero() { + *output = Self::zero(); + return; + } + + output.sign = if self.sign == other.sign { + Sign::Positive + } else { + Sign::Negative + }; + match (self.limbs.as_slice(), other.limbs.as_slice()) { + ([scalar], value) => mul_abs_one_limb_into(value, *scalar, &mut output.limbs), + (value, [scalar]) => mul_abs_one_limb_into(value, *scalar, &mut output.limbs), + _ => mul_abs_limbs_into(&self.limbs, &other.limbs, &mut output.limbs), + } + if output.limbs.is_empty() { + output.sign = Sign::Zero; + } + } + + pub(super) fn mul_i128(&self, scalar: i128) -> Self { + if self.is_zero() || scalar == 0 { + return Self::zero(); + } + if scalar == 1 { + return self.clone(); + } + if scalar == -1 { + return self.clone().negated(); + } + + let sign = if scalar < 0 { + match self.sign { + Sign::Negative => Sign::Positive, + Sign::Positive => Sign::Negative, + Sign::Zero => Sign::Zero, + } + } else { + self.sign + }; + Self::from_parts(sign, mul_abs_small(&self.limbs, scalar.unsigned_abs())) + } + + pub(super) fn linear_combination_i128_with_scratch( + left: &Self, + left_scalar: i128, + right: &Self, + right_scalar: i128, + scratch: &mut LimbScratch, + ) -> Self { + let mut output = Self::zero(); + Self::linear_combination_i128_into( + left, + left_scalar, + right, + right_scalar, + &mut output, + scratch, + ); + output + } + + #[inline(always)] + fn linear_combination_i128_into( + left: &Self, + left_scalar: i128, + right: &Self, + right_scalar: i128, + output: &mut Self, + scratch: &mut LimbScratch, + ) { + if left_scalar == 0 && right_scalar == 0 { + *output = Self::zero(); + return; + } + if left_scalar == 0 { + right.mul_i128_into(right_scalar, output); + return; + } + if right_scalar == 0 { + left.mul_i128_into(left_scalar, output); + return; + } + + let left_abs = left_scalar.unsigned_abs(); + let right_abs = right_scalar.unsigned_abs(); + let left_sign = signed_scalar_sign(left.sign, left_scalar); + let right_sign = signed_scalar_sign(right.sign, right_scalar); + + if left_abs == 1 && right_abs == 1 { + let sign = combine_signed_abs_limbs_into( + left_sign, + &left.limbs, + right_sign, + &right.limbs, + &mut output.limbs, + ); + output.sign = if output.limbs.is_empty() { + Sign::Zero + } else { + sign + }; + return; + } + + if left_abs == 1 { + mul_abs_small_into(&right.limbs, right_abs, &mut scratch.linear_right); + let sign = combine_signed_abs_limbs_into( + left_sign, + &left.limbs, + right_sign, + &scratch.linear_right, + &mut output.limbs, + ); + output.sign = if output.limbs.is_empty() { + Sign::Zero + } else { + sign + }; + return; + } + + if right_abs == 1 { + mul_abs_small_into(&left.limbs, left_abs, &mut scratch.linear_left); + let sign = combine_signed_abs_limbs_into( + left_sign, + &scratch.linear_left, + right_sign, + &right.limbs, + &mut output.limbs, + ); + output.sign = if output.limbs.is_empty() { + Sign::Zero + } else { + sign + }; + return; + } + + mul_abs_small_into(&left.limbs, left_abs, &mut scratch.linear_left); + mul_abs_small_into(&right.limbs, right_abs, &mut scratch.linear_right); + + let sign = combine_signed_abs_limbs_into( + left_sign, + &scratch.linear_left, + right_sign, + &scratch.linear_right, + &mut output.limbs, + ); + + output.sign = if output.limbs.is_empty() { + Sign::Zero + } else { + sign + }; + } + + #[inline(always)] + fn mul_i128_into(&self, scalar: i128, output: &mut Self) { + if self.is_zero() || scalar == 0 { + *output = Self::zero(); + return; + } + + let sign = if scalar < 0 { + match self.sign { + Sign::Negative => Sign::Positive, + Sign::Positive => Sign::Negative, + Sign::Zero => Sign::Zero, + } + } else { + self.sign + }; + output.sign = sign; + output.limbs.clear(); + mul_abs_small_into(&self.limbs, scalar.unsigned_abs(), &mut output.limbs); + if output.limbs.is_empty() { + output.sign = Sign::Zero; + } + } + + pub(super) fn square(&self) -> Self { + if self.is_zero() { + return Self::zero(); + } + Self::from_parts(Sign::Positive, square_abs_limbs(&self.limbs)) + } + + pub(super) fn square_into(&self, output: &mut Self) { + if self.is_zero() { + *output = Self::zero(); + return; + } + output.sign = Sign::Positive; + mul_abs_limbs_into(&self.limbs, &self.limbs, &mut output.limbs); + if output.limbs.is_empty() { + output.sign = Sign::Zero; + } + } + + pub(super) fn shl_bits(&self, bits: usize) -> Self { + if self.is_zero() || bits == 0 { + return self.clone(); + } + Self::from_parts(self.sign, shl_abs_limbs(&self.limbs, bits)) + } + + pub(super) fn shl_bits_into(&self, bits: usize, output: &mut Self) { + if self.is_zero() { + *output = Self::zero(); + return; + } + if bits == 0 { + *output = self.clone(); + return; + } + output.sign = self.sign; + shl_abs_limbs_into(&self.limbs, bits, &mut output.limbs); + if output.limbs.is_empty() { + output.sign = Sign::Zero; + } + } + + pub(super) fn shl_bits_owned(mut self, bits: usize) -> Self { + if self.is_zero() || bits == 0 { + return self; + } + shl_abs_limbs_assign(&mut self.limbs, bits); + self + } + + pub(super) fn shl_bits_assign(&mut self, bits: usize) { + if self.is_zero() || bits == 0 { + return; + } + shl_abs_limbs_assign(&mut self.limbs, bits); + } + + pub(super) fn shr_abs_bits(&self, bits: usize) -> Self { + if self.is_zero() || bits == 0 { + return self.abs(); + } + Self::from_parts(Sign::Positive, shr_abs_limbs(&self.limbs, bits)) + } + + pub(super) fn div_rem(&self, divisor: &Self) -> (Self, Self) { + let mut scratch = LimbScratch::default(); + self.div_rem_with_scratch(divisor, &mut scratch) + } + + pub(super) fn div_rem_with_scratch( + &self, + divisor: &Self, + scratch: &mut LimbScratch, + ) -> (Self, Self) { + assert!(!divisor.is_zero(), "division by zero"); + if self.is_zero() { + return (Self::zero(), Self::zero()); + } + + let (quotient_limbs, remainder_limbs) = + div_rem_abs_limbs_with_scratch(&self.limbs, &divisor.limbs, &mut scratch.division); + let quotient_sign = if quotient_limbs.is_empty() { + Sign::Zero + } else if self.sign == divisor.sign { + Sign::Positive + } else { + Sign::Negative + }; + let remainder_sign = if remainder_limbs.is_empty() { + Sign::Zero + } else { + self.sign + }; + + ( + Self::from_parts(quotient_sign, quotient_limbs), + Self::from_parts(remainder_sign, remainder_limbs), + ) + } + + pub(super) fn rem(&self, divisor: &Self) -> Self { + self.div_rem(divisor).1 + } + + pub(super) fn rem_with_scratch(&self, divisor: &Self, scratch: &mut LimbScratch) -> Self { + assert!(!divisor.is_zero(), "division by zero"); + if self.is_zero() { + return Self::zero(); + } + + let remainder_limbs = + rem_abs_limbs_with_scratch(&self.limbs, &divisor.limbs, &mut scratch.division); + let remainder_sign = if remainder_limbs.is_empty() { + Sign::Zero + } else { + self.sign + }; + Self::from_parts(remainder_sign, remainder_limbs) + } + + pub(super) fn div(&self, divisor: &Self) -> Self { + self.div_rem(divisor).0 + } + + pub(super) fn div_with_scratch(&self, divisor: &Self, scratch: &mut LimbScratch) -> Self { + assert!(!divisor.is_zero(), "division by zero"); + if self.is_zero() { + return Self::zero(); + } + match self.abs_cmp(divisor) { + Ordering::Less => return Self::zero(), + Ordering::Equal => { + return if self.sign == divisor.sign { + Self::one() + } else { + Self::from_i128(-1) + }; + } + Ordering::Greater => { + if cmp_abs_to_double(&self.limbs, &divisor.limbs) == Ordering::Less { + return if self.sign == divisor.sign { + Self::one() + } else { + Self::from_i128(-1) + }; + } + } + } + let quotient_limbs = + div_abs_limbs_with_scratch(&self.limbs, &divisor.limbs, &mut scratch.division); + let quotient_sign = if quotient_limbs.is_empty() { + Sign::Zero + } else if self.sign == divisor.sign { + Sign::Positive + } else { + Sign::Negative + }; + Self::from_parts(quotient_sign, quotient_limbs) + } + + pub(super) fn div2_exact(&self) -> Self { + debug_assert!(self.limbs.first().copied().unwrap_or(0) & 1 == 0); + if self.is_zero() { + return Self::zero(); + } + Self::from_parts(self.sign, shr_abs_limbs(&self.limbs, 1)) + } + + pub(super) fn div_floor(&self, divisor: &Self) -> Self { + let (quotient, remainder) = self.div_rem(divisor); + self.finish_div_floor(divisor, quotient, remainder) + } + + pub(super) fn div_floor_with_scratch(&self, divisor: &Self, scratch: &mut LimbScratch) -> Self { + assert!(!divisor.is_zero(), "division by zero"); + if self.is_zero() { + return Self::zero(); + } + + match self.abs_cmp(divisor) { + Ordering::Less => { + return if self.sign == divisor.sign { + Self::zero() + } else { + Self::from_i128(-1) + }; + } + Ordering::Equal => { + return if self.sign == divisor.sign { + Self::one() + } else { + Self::from_i128(-1) + }; + } + Ordering::Greater => { + if cmp_abs_to_double(&self.limbs, &divisor.limbs) == Ordering::Less { + return if self.sign == divisor.sign { + Self::one() + } else { + Self::from_i128(-2) + }; + } + } + } + + if self.sign == divisor.sign { + let quotient_limbs = + div_abs_limbs_with_scratch(&self.limbs, &divisor.limbs, &mut scratch.division); + let quotient_sign = if quotient_limbs.is_empty() { + Sign::Zero + } else { + Sign::Positive + }; + return Self::from_parts(quotient_sign, quotient_limbs); + } + + let (quotient, remainder) = self.div_rem_with_scratch(divisor, scratch); + self.finish_div_floor(divisor, quotient, remainder) + } + + fn finish_div_floor(&self, divisor: &Self, quotient: Self, remainder: Self) -> Self { + if remainder.is_zero() || self.sign == divisor.sign { + quotient + } else { + quotient.sub(&Self::one()) + } + } + + pub(super) fn mod_positive(&self, modulus: &Self) -> Self { + let mut scratch = LimbScratch::default(); + self.mod_positive_with_scratch(modulus, &mut scratch) + } + + pub(super) fn mod_positive_with_scratch( + &self, + modulus: &Self, + scratch: &mut LimbScratch, + ) -> Self { + assert!(!modulus.is_zero(), "division by zero"); + if self.is_zero() { + return Self::zero(); + } + match self.abs_cmp(modulus) { + Ordering::Less => { + return if self.is_negative() { + self.add(modulus) + } else { + self.clone() + }; + } + Ordering::Equal | Ordering::Greater => {} + } + + let remainder_limbs = + rem_abs_limbs_with_scratch(&self.limbs, &modulus.limbs, &mut scratch.division); + let remainder_sign = if remainder_limbs.is_empty() { + Sign::Zero + } else { + self.sign + }; + let mut remainder = Self::from_parts(remainder_sign, remainder_limbs); + if remainder.is_negative() { + remainder = remainder.add(modulus); + } + remainder + } + + pub(super) fn abs(&self) -> Self { + if self.is_zero() { + Self::zero() + } else { + Self { + sign: Sign::Positive, + limbs: self.limbs.clone(), + } + } + } + + #[inline(always)] + pub(super) fn shifted_low_word(&self, shift_bits: u64) -> u64 { + debug_assert!(!self.is_negative()); + let digit_index = usize::try_from(shift_bits / 64).unwrap_or(usize::MAX); + let offset = (shift_bits % 64) as u32; + let low = self.limbs.get(digit_index).copied().unwrap_or(0); + if offset == 0 { + return low; + } + let high = self.limbs.get(digit_index + 1).copied().unwrap_or(0); + (low >> offset) | (high << (64 - offset)) + } + + pub(super) fn extended_gcd(&self, other: &Self) -> ExtendedGcd { + let mut scratch = LimbScratch::default(); + self.extended_gcd_with_scratch(other, &mut scratch) + } + + pub(super) fn extended_gcd_with_scratch( + &self, + other: &Self, + scratch: &mut LimbScratch, + ) -> ExtendedGcd { + let mut old_r = self.clone(); + let mut r = other.clone(); + let mut old_s = Self::one(); + let mut s = Self::zero(); + let mut old_t = Self::zero(); + let mut t = Self::one(); + let mut next_old_r = Self::zero(); + let mut next_r = Self::zero(); + let mut next_old_s = Self::zero(); + let mut next_s = Self::zero(); + let mut next_old_t = Self::zero(); + let mut next_t = Self::zero(); + + while !r.is_zero() { + if old_r.is_one() { + break; + } + if r.is_one() { + old_r = r; + old_s = s; + old_t = t; + break; + } + + if !old_r.is_negative() && !r.is_negative() && old_r.abs_cmp(&r) == Ordering::Less { + std::mem::swap(&mut old_r, &mut r); + std::mem::swap(&mut old_s, &mut s); + std::mem::swap(&mut old_t, &mut t); + continue; + } + + if !old_r.is_negative() && !r.is_negative() { + let bits = old_r.bit_len().saturating_sub(63); + let mut rr2 = old_r.shifted_low_word(bits); + let mut rr1 = r.shifted_low_word(bits); + + let mut aa2 = 0_i128; + let mut aa1 = 1_i128; + let mut bb2 = 1_i128; + let mut bb1 = 0_i128; + let mut steps = 0_u32; + + while rr1 != 0 { + let q = rr2 / rr1; + if q == 0 { + break; + } + let next_r = rr2 - q * rr1; + let q = i128::from(q); + let next_a = aa2 - q * aa1; + let next_b = bb2 - q * bb1; + let next_r_i = i128::from(next_r); + let rr1_minus_next_r = i128::from(rr1 - next_r); + + if steps & 1 == 1 { + if next_r_i < -next_b || rr1_minus_next_r < next_a - aa1 { + break; + } + } else if next_r_i < -next_a || rr1_minus_next_r < next_b - bb1 { + break; + } + + rr2 = rr1; + rr1 = next_r; + aa2 = aa1; + aa1 = next_a; + bb2 = bb1; + bb1 = next_b; + steps += 1; + } + + if steps != 0 { + Self::linear_combination_i128_into( + &old_r, + bb2, + &r, + aa2, + &mut next_old_r, + scratch, + ); + Self::linear_combination_i128_into(&r, aa1, &old_r, bb1, &mut next_r, scratch); + Self::linear_combination_i128_into( + &old_s, + bb2, + &s, + aa2, + &mut next_old_s, + scratch, + ); + Self::linear_combination_i128_into(&s, aa1, &old_s, bb1, &mut next_s, scratch); + Self::linear_combination_i128_into( + &old_t, + bb2, + &t, + aa2, + &mut next_old_t, + scratch, + ); + Self::linear_combination_i128_into(&t, aa1, &old_t, bb1, &mut next_t, scratch); + + std::mem::swap(&mut old_r, &mut next_old_r); + std::mem::swap(&mut r, &mut next_r); + std::mem::swap(&mut old_s, &mut next_old_s); + std::mem::swap(&mut s, &mut next_s); + std::mem::swap(&mut old_t, &mut next_old_t); + std::mem::swap(&mut t, &mut next_t); + + if old_r.is_negative() { + old_r = old_r.negated(); + old_s = old_s.negated(); + old_t = old_t.negated(); + } + if r.is_negative() { + r = r.negated(); + s = s.negated(); + t = t.negated(); + } + continue; + } + } + + let (quotient, next_r) = old_r.div_rem_with_scratch(&r, scratch); + old_r = r; + r = next_r; + + quotient.mul_into(&s, &mut next_s); + let mut fallback_next_s = std::mem::take(&mut old_s); + fallback_next_s.sub_assign(&next_s); + old_s = s; + s = fallback_next_s; + + quotient.mul_into(&t, &mut next_t); + let mut fallback_next_t = std::mem::take(&mut old_t); + fallback_next_t.sub_assign(&next_t); + old_t = t; + t = fallback_next_t; + } + + if old_r.is_negative() { + old_r = old_r.negated(); + old_s = old_s.negated(); + old_t = old_t.negated(); + } + + ExtendedGcd { + x: old_s, + y: old_t, + gcd: old_r, + } + } + + pub(super) fn left_extended_gcd_with_scratch( + &self, + other: &Self, + scratch: &mut LimbScratch, + ) -> LeftExtendedGcd { + let mut old_r = self.clone(); + let mut r = other.clone(); + let mut old_s = Self::one(); + let mut s = Self::zero(); + let mut next_old_r = Self::zero(); + let mut next_r = Self::zero(); + let mut next_old_s = Self::zero(); + let mut next_s = Self::zero(); + + while !r.is_zero() { + if old_r.is_one() { + break; + } + if r.is_one() { + old_r = r; + old_s = s; + break; + } + + if !old_r.is_negative() && !r.is_negative() && old_r.abs_cmp(&r) == Ordering::Less { + std::mem::swap(&mut old_r, &mut r); + std::mem::swap(&mut old_s, &mut s); + continue; + } + + if !old_r.is_negative() && !r.is_negative() { + let bits = old_r.bit_len().saturating_sub(63); + let mut rr2 = old_r.shifted_low_word(bits); + let mut rr1 = r.shifted_low_word(bits); + + let mut aa2 = 0_i128; + let mut aa1 = 1_i128; + let mut bb2 = 1_i128; + let mut bb1 = 0_i128; + let mut steps = 0_u32; + + while rr1 != 0 { + let q = rr2 / rr1; + if q == 0 { + break; + } + let next_r = rr2 - q * rr1; + let q = i128::from(q); + let next_a = aa2 - q * aa1; + let next_b = bb2 - q * bb1; + let next_r_i = i128::from(next_r); + let rr1_minus_next_r = i128::from(rr1 - next_r); + + if steps & 1 == 1 { + if next_r_i < -next_b || rr1_minus_next_r < next_a - aa1 { + break; + } + } else if next_r_i < -next_a || rr1_minus_next_r < next_b - bb1 { + break; + } + + rr2 = rr1; + rr1 = next_r; + aa2 = aa1; + aa1 = next_a; + bb2 = bb1; + bb1 = next_b; + steps += 1; + } + + if steps != 0 { + Self::linear_combination_i128_into( + &old_r, + bb2, + &r, + aa2, + &mut next_old_r, + scratch, + ); + Self::linear_combination_i128_into(&r, aa1, &old_r, bb1, &mut next_r, scratch); + Self::linear_combination_i128_into( + &old_s, + bb2, + &s, + aa2, + &mut next_old_s, + scratch, + ); + Self::linear_combination_i128_into(&s, aa1, &old_s, bb1, &mut next_s, scratch); + + std::mem::swap(&mut old_r, &mut next_old_r); + std::mem::swap(&mut r, &mut next_r); + std::mem::swap(&mut old_s, &mut next_old_s); + std::mem::swap(&mut s, &mut next_s); + + if old_r.is_negative() { + old_r = old_r.negated(); + old_s = old_s.negated(); + } + if r.is_negative() { + r = r.negated(); + s = s.negated(); + } + continue; + } + } + + let (quotient, next_r) = old_r.div_rem_with_scratch(&r, scratch); + old_r = r; + r = next_r; + + quotient.mul_into(&s, &mut next_s); + let mut fallback_next_s = std::mem::take(&mut old_s); + fallback_next_s.sub_assign(&next_s); + old_s = s; + s = fallback_next_s; + } + + if old_r.is_negative() { + old_r = old_r.negated(); + old_s = old_s.negated(); + } + + LeftExtendedGcd { + x: old_s, + gcd: old_r, + } + } + + pub(super) fn left_extended_gcd_positive_with_scratch( + &self, + other: &Self, + scratch: &mut LimbScratch, + ) -> LeftExtendedGcd { + debug_assert!(!self.is_negative()); + debug_assert!(!other.is_negative()); + + let mut old_r; + let mut r; + let mut old_s; + let mut s; + if self.abs_cmp(other) == Ordering::Less { + old_r = other.clone(); + r = self.clone(); + old_s = Self::zero(); + s = Self::one(); + } else { + old_r = self.clone(); + r = other.clone(); + old_s = Self::one(); + s = Self::zero(); + } + let mut next_old_r = Self::zero(); + let mut next_r = Self::zero(); + let mut next_old_s = Self::zero(); + let mut next_s = Self::zero(); + + while !r.is_zero() { + if old_r.is_one() { + break; + } + if r.is_one() { + old_r = r; + old_s = s; + break; + } + + let bits = old_r.bit_len().saturating_sub(63); + let mut rr2 = old_r.shifted_low_word(bits); + let mut rr1 = r.shifted_low_word(bits); + + let mut aa2 = 0_i128; + let mut aa1 = 1_i128; + let mut bb2 = 1_i128; + let mut bb1 = 0_i128; + let mut steps = 0_u32; + + while rr1 != 0 { + let q = rr2 / rr1; + if q == 0 { + break; + } + let next_r = rr2 - q * rr1; + let q = i128::from(q); + let next_a = aa2 - q * aa1; + let next_b = bb2 - q * bb1; + let next_r_i = i128::from(next_r); + let rr1_minus_next_r = i128::from(rr1 - next_r); + + if steps & 1 == 1 { + if next_r_i < -next_b || rr1_minus_next_r < next_a - aa1 { + break; + } + } else if next_r_i < -next_a || rr1_minus_next_r < next_b - bb1 { + break; + } + + rr2 = rr1; + rr1 = next_r; + aa2 = aa1; + aa1 = next_a; + bb2 = bb1; + bb1 = next_b; + steps += 1; + } + + if steps != 0 { + Self::linear_combination_i128_into(&old_r, bb2, &r, aa2, &mut next_old_r, scratch); + Self::linear_combination_i128_into(&r, aa1, &old_r, bb1, &mut next_r, scratch); + Self::linear_combination_i128_into(&old_s, bb2, &s, aa2, &mut next_old_s, scratch); + Self::linear_combination_i128_into(&s, aa1, &old_s, bb1, &mut next_s, scratch); + + std::mem::swap(&mut old_r, &mut next_old_r); + std::mem::swap(&mut r, &mut next_r); + std::mem::swap(&mut old_s, &mut next_old_s); + std::mem::swap(&mut s, &mut next_s); + + if old_r.is_negative() { + old_r = old_r.negated(); + old_s = old_s.negated(); + } + if r.is_negative() { + r = r.negated(); + s = s.negated(); + } + continue; + } + + let (quotient, next_r) = old_r.div_rem_with_scratch(&r, scratch); + old_r = r; + r = next_r; + + quotient.mul_into(&s, &mut next_s); + let mut fallback_next_s = std::mem::take(&mut old_s); + fallback_next_s.sub_assign(&next_s); + old_s = s; + s = fallback_next_s; + } + + if old_r.is_negative() { + old_r = old_r.negated(); + old_s = old_s.negated(); + } + + LeftExtendedGcd { + x: old_s, + gcd: old_r, + } + } + + fn from_parts(sign: Sign, mut limbs: LimbVec) -> Self { + trim_leading_zero_limbs(&mut limbs); + let sign = if limbs.is_empty() { Sign::Zero } else { sign }; + Self { sign, limbs } + } + + fn abs_biguint(&self) -> num_bigint::BigUint { + let mut bytes = Vec::with_capacity(self.limbs.len() * 8); + for limb in &self.limbs { + bytes.extend_from_slice(&limb.to_le_bytes()); + } + num_bigint::BigUint::from_bytes_le(&bytes) + } +} + +pub(super) fn xgcd_partial( + r2: &LimbInt, + r1: &LimbInt, + threshold: &LimbInt, +) -> (LimbInt, LimbInt, LimbInt, LimbInt) { + let mut scratch = LimbScratch::default(); + xgcd_partial_with_scratch(r2, r1, threshold, &mut scratch) +} + +pub(super) fn xgcd_partial_with_scratch( + r2: &LimbInt, + r1: &LimbInt, + threshold: &LimbInt, + scratch: &mut LimbScratch, +) -> (LimbInt, LimbInt, LimbInt, LimbInt) { + let mut r2 = r2.clone(); + let mut r1 = r1.clone(); + let mut co2 = LimbInt::zero(); + let mut co1 = LimbInt::from_i128(-1); + let mut next_r2 = LimbInt::zero(); + let mut next_r1 = LimbInt::zero(); + let mut next_co2 = LimbInt::zero(); + let mut next_co1 = LimbInt::zero(); + + while !r1.is_zero() && r1.cmp(threshold) == Ordering::Greater { + let bits = r2.bit_len().saturating_sub(63); + let mut rr2 = r2.shifted_low_word(bits); + let mut rr1 = r1.shifted_low_word(bits); + let threshold_word = threshold.shifted_low_word(bits); + + let mut aa2 = 0_i128; + let mut aa1 = 1_i128; + let mut bb2 = 1_i128; + let mut bb1 = 0_i128; + let mut steps = 0_u32; + + while rr1 != 0 && rr1 > threshold_word { + let q = rr2 / rr1; + if q == 0 { + break; + } + let next_r = rr2 - q * rr1; + let q = i128::from(q); + let next_a = aa2 - q * aa1; + let next_b = bb2 - q * bb1; + let next_r_i = i128::from(next_r); + let rr1_minus_next_r = i128::from(rr1 - next_r); + + if steps & 1 == 1 { + if next_r_i < -next_b || rr1_minus_next_r < next_a - aa1 { + break; + } + } else if next_r_i < -next_a || rr1_minus_next_r < next_b - bb1 { + break; + } + + rr2 = rr1; + rr1 = next_r; + aa2 = aa1; + aa1 = next_a; + bb2 = bb1; + bb1 = next_b; + steps += 1; + } + + if steps == 0 { + let (q, next_r) = r2.div_rem_with_scratch(&r1, scratch); + q.mul_into(&co1, &mut next_co1); + let mut next_co = std::mem::take(&mut co2); + next_co.sub_assign(&next_co1); + r2 = r1; + r1 = next_r; + co2 = co1; + co1 = next_co; + } else { + LimbInt::linear_combination_i128_into(&r2, bb2, &r1, aa2, &mut next_r2, scratch); + LimbInt::linear_combination_i128_into(&r1, aa1, &r2, bb1, &mut next_r1, scratch); + LimbInt::linear_combination_i128_into(&co2, bb2, &co1, aa2, &mut next_co2, scratch); + LimbInt::linear_combination_i128_into(&co1, aa1, &co2, bb1, &mut next_co1, scratch); + + std::mem::swap(&mut r2, &mut next_r2); + std::mem::swap(&mut r1, &mut next_r1); + std::mem::swap(&mut co2, &mut next_co2); + std::mem::swap(&mut co1, &mut next_co1); + + if r1.is_negative() { + r1 = r1.negated(); + co1 = co1.negated(); + } + if r2.is_negative() { + r2 = r2.negated(); + co2 = co2.negated(); + } + } + } + + if r2.is_negative() { + r2 = r2.negated(); + co2 = co2.negated(); + co1 = co1.negated(); + } + + (co2, co1, r2, r1) +} + +fn subtract_signed_abs( + positive_limbs: &[u64], + positive_sign: Sign, + negative_limbs: &[u64], + negative_sign: Sign, +) -> LimbInt { + match cmp_abs_limbs(positive_limbs, negative_limbs) { + Ordering::Greater => { + LimbInt::from_parts(positive_sign, sub_abs_limbs(positive_limbs, negative_limbs)) + } + Ordering::Less => { + LimbInt::from_parts(negative_sign, sub_abs_limbs(negative_limbs, positive_limbs)) + } + Ordering::Equal => LimbInt::zero(), + } +} + +fn signed_scalar_sign(value_sign: Sign, scalar: i128) -> Sign { + if scalar == 0 || value_sign == Sign::Zero { + Sign::Zero + } else if scalar < 0 { + match value_sign { + Sign::Negative => Sign::Positive, + Sign::Zero => Sign::Zero, + Sign::Positive => Sign::Negative, + } + } else { + value_sign + } +} + +#[inline(always)] +fn combine_signed_abs_limbs_into( + left_sign: Sign, + left: &[u64], + right_sign: Sign, + right: &[u64], + output: &mut LimbVec, +) -> Sign { + output.clear(); + match (left_sign, right_sign) { + (Sign::Zero, Sign::Zero) => Sign::Zero, + (Sign::Zero, _) => { + output.extend_from_slice(right); + right_sign + } + (_, Sign::Zero) => { + output.extend_from_slice(left); + left_sign + } + (Sign::Positive, Sign::Positive) | (Sign::Negative, Sign::Negative) => { + add_abs_limbs_into(left, right, output); + left_sign + } + (Sign::Positive, Sign::Negative) | (Sign::Negative, Sign::Positive) => { + match cmp_abs_limbs(left, right) { + Ordering::Greater => { + sub_abs_limbs_into(left, right, output); + left_sign + } + Ordering::Less => { + sub_abs_limbs_into(right, left, output); + right_sign + } + Ordering::Equal => Sign::Zero, + } + } + } +} + +#[inline(always)] +fn cmp_abs_limbs(left: &[u64], right: &[u64]) -> Ordering { + match left.len().cmp(&right.len()) { + Ordering::Equal => left.iter().rev().cmp(right.iter().rev()), + other => other, + } +} + +fn cmp_abs_to_double(left: &[u64], right: &[u64]) -> Ordering { + debug_assert!(!right.is_empty()); + let doubled_len = right.len() + usize::from(right.last().copied().unwrap_or(0) >> 63 != 0); + match left.len().cmp(&doubled_len) { + Ordering::Equal => {} + other => return other, + } + + for index in (0..doubled_len).rev() { + let doubled_limb = shifted_left_one_limb_unchecked(right, index); + match left[index].cmp(&doubled_limb) { + Ordering::Equal => {} + other => return other, + } + } + Ordering::Equal +} + +fn shifted_left_one_limb_unchecked(limbs: &[u64], index: usize) -> u64 { + debug_assert!(index <= limbs.len()); + let low = if index < limbs.len() { + limbs[index] << 1 + } else { + 0 + }; + let carry = if index == 0 { + 0 + } else { + limbs[index - 1] >> 63 + }; + low | carry +} + +fn add_abs_limbs(left: &[u64], right: &[u64]) -> LimbVec { + let mut output = LimbVec::new(); + add_abs_limbs_into(left, right, &mut output); + output +} + +#[inline(always)] +#[allow(clippy::uninit_vec)] +fn add_abs_limbs_into(left: &[u64], right: &[u64], output: &mut LimbVec) { + output.clear(); + let len = left.len().max(right.len()); + output.reserve(len + 1); + // SAFETY: u64 has no drop glue, reserve guarantees room for len + 1 limbs, + // and the loops initialize every slot in 0..len before the vector is read. + unsafe { + output.set_len(len); + } + let mut carry = 0_u64; + let shared_len = left.len().min(right.len()); + + for index in 0..shared_len { + let (sum, carry_a) = left[index].overflowing_add(right[index]); + let (sum, carry_b) = sum.overflowing_add(carry); + output[index] = sum; + carry = u64::from(carry_a || carry_b); + } + + let remaining = if left.len() > right.len() { + &left[shared_len..] + } else { + &right[shared_len..] + }; + for (offset, limb) in remaining.iter().copied().enumerate() { + let (sum, next_carry) = limb.overflowing_add(carry); + output[shared_len + offset] = sum; + carry = u64::from(next_carry); + } + if carry != 0 { + output.push(carry); + } +} + +fn add_abs_limbs_assign(left: &mut LimbVec, right: &[u64]) { + let len = left.len().max(right.len()); + let original_left_len = left.len(); + left.resize(len, 0); + let mut carry = 0_u64; + let shared_len = original_left_len.min(right.len()); + + for index in 0..shared_len { + let (sum, carry_a) = left[index].overflowing_add(right[index]); + let (sum, carry_b) = sum.overflowing_add(carry); + left[index] = sum; + carry = u64::from(carry_a || carry_b); + } + + if right.len() > original_left_len { + for index in shared_len..right.len() { + let (sum, next_carry) = right[index].overflowing_add(carry); + left[index] = sum; + carry = u64::from(next_carry); + } + } else { + for left_limb in left.iter_mut().take(original_left_len).skip(shared_len) { + let (sum, next_carry) = left_limb.overflowing_add(carry); + *left_limb = sum; + carry = u64::from(next_carry); + } + } + + if carry != 0 { + left.push(carry); + } +} + +fn sub_abs_limbs(left: &[u64], right: &[u64]) -> LimbVec { + let mut output = LimbVec::new(); + sub_abs_limbs_into(left, right, &mut output); + output +} + +#[inline(always)] +#[allow(clippy::uninit_vec)] +fn sub_abs_limbs_into(left: &[u64], right: &[u64], output: &mut LimbVec) { + debug_assert!(cmp_abs_limbs(left, right) != Ordering::Less); + output.clear(); + output.reserve(left.len()); + // SAFETY: u64 has no drop glue, reserve guarantees room for left.len() + // limbs, and the loops initialize every slot before trimming reads them. + unsafe { + output.set_len(left.len()); + } + let mut borrow = 0_u64; + let shared_len = right.len(); + + for index in 0..shared_len { + let (difference, borrow_a) = left[index].overflowing_sub(right[index]); + let (difference, borrow_b) = difference.overflowing_sub(borrow); + output[index] = difference; + borrow = u64::from(borrow_a || borrow_b); + } + + for (offset, left_limb) in left[shared_len..].iter().copied().enumerate() { + let (difference, next_borrow) = left_limb.overflowing_sub(borrow); + output[shared_len + offset] = difference; + borrow = u64::from(next_borrow); + } + + debug_assert_eq!(borrow, 0); + trim_leading_zero_limbs(output); +} + +fn sub_abs_limbs_assign(left: &mut LimbVec, right: &[u64]) { + debug_assert!(cmp_abs_limbs(left, right) != Ordering::Less); + let mut borrow = 0_u64; + let shared_len = right.len(); + + for index in 0..shared_len { + let (difference, borrow_a) = left[index].overflowing_sub(right[index]); + let (difference, borrow_b) = difference.overflowing_sub(borrow); + left[index] = difference; + borrow = u64::from(borrow_a || borrow_b); + } + + let mut index = shared_len; + while borrow != 0 && index < left.len() { + let (difference, next_borrow) = left[index].overflowing_sub(borrow); + left[index] = difference; + borrow = u64::from(next_borrow); + index += 1; + } + + debug_assert_eq!(borrow, 0); + trim_leading_zero_limbs(left); +} + +fn mul_abs_limbs(left: &[u64], right: &[u64]) -> LimbVec { + let mut output = LimbVec::new(); + mul_abs_limbs_into(left, right, &mut output); + output +} + +fn mul_abs_limbs_into(left: &[u64], right: &[u64], output: &mut LimbVec) { + output.clear(); + if left.is_empty() || right.is_empty() { + return; + } + + let (outer, inner) = if left.len() <= right.len() { + (left, right) + } else { + (right, left) + }; + + output.resize(outer.len() + inner.len(), 0); + for (left_index, left_limb) in outer.iter().copied().enumerate() { + let mut carry = 0_u128; + for (right_index, right_limb) in inner.iter().copied().enumerate() { + let output_index = left_index + right_index; + let product = u128::from(left_limb) * u128::from(right_limb) + + u128::from(output[output_index]) + + carry; + output[output_index] = product as u64; + carry = product >> 64; + } + + let mut output_index = left_index + inner.len(); + while carry != 0 { + let sum = u128::from(output[output_index]) + carry; + output[output_index] = sum as u64; + carry = sum >> 64; + output_index += 1; + } + } + + trim_leading_zero_limbs(output); +} + +fn mul_abs_one_limb(value: &[u64], scalar: u64) -> LimbVec { + let mut output = LimbVec::new(); + mul_abs_one_limb_into(value, scalar, &mut output); + output +} + +#[inline(always)] +#[allow(clippy::uninit_vec)] +fn mul_abs_one_limb_into(value: &[u64], scalar: u64, output: &mut LimbVec) { + output.clear(); + if value.is_empty() || scalar == 0 { + return; + } + if scalar == 1 { + output.extend_from_slice(value); + return; + } + + output.reserve(value.len() + 1); + let mut carry = 0_u128; + // SAFETY: u64 has no drop glue, reserve guarantees room for value.len() + 1 + // limbs, and the loop initializes every slot before the vector is read. + unsafe { + output.set_len(value.len() + 1); + } + for (index, limb) in value.iter().copied().enumerate() { + let product = u128::from(limb) * u128::from(scalar) + carry; + output[index] = product as u64; + carry = product >> 64; + } + if carry != 0 { + output[value.len()] = carry as u64; + } else { + output.truncate(value.len()); + } +} + +fn square_abs_limbs(value: &[u64]) -> LimbVec { + mul_abs_limbs(value, value) +} + +fn mul_abs_small(value: &[u64], scalar: u128) -> LimbVec { + let mut output = LimbVec::new(); + mul_abs_small_into(value, scalar, &mut output); + output +} + +#[inline(always)] +fn mul_abs_small_into(value: &[u64], scalar: u128, output: &mut LimbVec) { + output.clear(); + if value.is_empty() || scalar == 0 { + return; + } + if scalar == 1 { + output.extend_from_slice(value); + return; + } + let scalar_low = scalar as u64; + let scalar_high = (scalar >> 64) as u64; + if scalar_high == 0 { + mul_abs_one_limb_into(value, scalar_low, output); + return; + } + + output.resize(value.len() + 2, 0); + + if scalar_low != 0 { + for (index, limb) in value.iter().copied().enumerate() { + add_u128_at(output, index, u128::from(limb) * u128::from(scalar_low)); + } + } + if scalar_high != 0 { + for (index, limb) in value.iter().copied().enumerate() { + add_u128_at( + output, + index + 1, + u128::from(limb) * u128::from(scalar_high), + ); + } + } + + trim_leading_zero_limbs(output); +} + +fn add_u128_at(output: &mut LimbVec, index: usize, value: u128) { + let low = value as u64; + let high = (value >> 64) as u64; + ensure_len(output, index + 2); + + let (sum_low, carry_low) = output[index].overflowing_add(low); + output[index] = sum_low; + + let (sum_high, carry_high_a) = output[index + 1].overflowing_add(high); + let (sum_high, carry_high_b) = sum_high.overflowing_add(u64::from(carry_low)); + output[index + 1] = sum_high; + + let mut carry = u64::from(carry_high_a) + u64::from(carry_high_b); + let mut carry_index = index + 2; + while carry != 0 { + ensure_len(output, carry_index + 1); + let (sum, overflowed) = output[carry_index].overflowing_add(carry); + output[carry_index] = sum; + carry = u64::from(overflowed); + carry_index += 1; + } +} + +fn ensure_len(output: &mut LimbVec, len: usize) { + if output.len() < len { + output.resize(len, 0); + } +} + +fn div_rem_abs_limbs(numerator: &[u64], denominator: &[u64]) -> (LimbVec, LimbVec) { + let mut scratch = DivisionScratch::default(); + div_rem_abs_limbs_with_scratch(numerator, denominator, &mut scratch) +} + +fn div_rem_abs_limbs_with_scratch( + numerator: &[u64], + denominator: &[u64], + scratch: &mut DivisionScratch, +) -> (LimbVec, LimbVec) { + div_rem_abs_limbs_into_scratch(numerator, denominator, scratch, true, true); + (scratch.quotient.clone(), scratch.remainder.clone()) +} + +fn div_abs_limbs_with_scratch( + numerator: &[u64], + denominator: &[u64], + scratch: &mut DivisionScratch, +) -> LimbVec { + div_rem_abs_limbs_into_scratch(numerator, denominator, scratch, true, false); + scratch.quotient.clone() +} + +fn rem_abs_limbs_with_scratch( + numerator: &[u64], + denominator: &[u64], + scratch: &mut DivisionScratch, +) -> LimbVec { + div_rem_abs_limbs_into_scratch(numerator, denominator, scratch, false, true); + scratch.remainder.clone() +} + +#[inline(always)] +fn div_rem_abs_limbs_into_scratch( + numerator: &[u64], + denominator: &[u64], + scratch: &mut DivisionScratch, + keep_quotient: bool, + keep_remainder: bool, +) { + debug_assert!(!denominator.is_empty()); + scratch.quotient.clear(); + scratch.remainder.clear(); + + if numerator.is_empty() { + return; + } + if cmp_abs_limbs(numerator, denominator) == Ordering::Less { + if keep_remainder { + scratch.remainder.extend_from_slice(numerator); + } + return; + } + if denominator == [1] { + if keep_quotient { + scratch.quotient.extend_from_slice(numerator); + } + return; + } + if denominator.len() == 1 { + div_rem_abs_one_limb_into( + numerator, + denominator[0], + &mut scratch.quotient, + &mut scratch.remainder, + keep_quotient, + keep_remainder, + ); + return; + } + + let shift = denominator.last().copied().unwrap_or(0).leading_zeros() as usize; + shl_abs_limbs_into(numerator, shift, &mut scratch.normalized_numerator); + shl_abs_limbs_into(denominator, shift, &mut scratch.normalized_denominator); + scratch.normalized_numerator.push(0); + + let denominator_len = scratch.normalized_denominator.len(); + let quotient_len = scratch.normalized_numerator.len() - denominator_len; + if keep_quotient { + scratch.quotient.resize(quotient_len, 0); + } + + for quotient_index in (0..quotient_len).rev() { + let (mut qhat, mut rhat, mut rhat_overflowed) = estimate_quotient( + scratch.normalized_numerator[quotient_index + denominator_len], + scratch.normalized_numerator[quotient_index + denominator_len - 1], + scratch.normalized_denominator[denominator_len - 1], + ); + + if denominator_len > 1 { + let next_denominator_limb = scratch.normalized_denominator[denominator_len - 2]; + let next_numerator_limb = + scratch.normalized_numerator[quotient_index + denominator_len - 2]; + while !rhat_overflowed + && quotient_too_large(qhat, rhat, next_denominator_limb, next_numerator_limb) + { + qhat -= 1; + let (next_rhat, overflowed) = + rhat.overflowing_add(scratch.normalized_denominator[denominator_len - 1]); + rhat = next_rhat; + rhat_overflowed = overflowed; + } + } + + if sub_mul_at( + &mut scratch.normalized_numerator, + &scratch.normalized_denominator, + qhat, + quotient_index, + ) { + qhat -= 1; + add_at( + &mut scratch.normalized_numerator, + &scratch.normalized_denominator, + quotient_index, + ); + } + if keep_quotient { + scratch.quotient[quotient_index] = qhat; + } + } + + if keep_remainder { + shr_abs_limbs_into( + &scratch.normalized_numerator[..denominator_len], + shift, + &mut scratch.remainder, + ); + trim_leading_zero_limbs(&mut scratch.remainder); + } + if keep_quotient { + trim_leading_zero_limbs(&mut scratch.quotient); + } +} + +fn div_rem_abs_one_limb(numerator: &[u64], denominator: u64) -> (LimbVec, LimbVec) { + let mut quotient = LimbVec::new(); + let mut remainder = LimbVec::new(); + div_rem_abs_one_limb_into( + numerator, + denominator, + &mut quotient, + &mut remainder, + true, + true, + ); + (quotient, remainder) +} + +fn div_rem_abs_one_limb_into( + numerator: &[u64], + denominator: u64, + quotient: &mut LimbVec, + remainder_output: &mut LimbVec, + keep_quotient: bool, + keep_remainder: bool, +) { + debug_assert_ne!(denominator, 0); + quotient.clear(); + remainder_output.clear(); + if keep_quotient { + quotient.resize(numerator.len(), 0); + } + let mut remainder = 0_u128; + + for (index, limb) in numerator.iter().copied().enumerate().rev() { + let value = (remainder << 64) | u128::from(limb); + if keep_quotient { + quotient[index] = (value / u128::from(denominator)) as u64; + } + remainder = value % u128::from(denominator); + } + + if keep_quotient { + trim_leading_zero_limbs(quotient); + } + if keep_remainder && remainder != 0 { + remainder_output.push(remainder as u64); + } +} + +#[inline(always)] +fn estimate_quotient(high: u64, low: u64, denominator_high: u64) -> (u64, u64, bool) { + if high == denominator_high { + let (remainder, overflowed) = low.overflowing_add(denominator_high); + (u64::MAX, remainder, overflowed) + } else { + let numerator = (u128::from(high) << 64) | u128::from(low); + ( + (numerator / u128::from(denominator_high)) as u64, + (numerator % u128::from(denominator_high)) as u64, + false, + ) + } +} + +#[inline(always)] +fn quotient_too_large(qhat: u64, rhat: u64, denominator_next: u64, numerator_next: u64) -> bool { + let left = u128::from(qhat) * u128::from(denominator_next); + let right = (u128::from(rhat) << 64) | u128::from(numerator_next); + left > right +} + +#[inline(always)] +fn sub_mul_at(target: &mut [u64], value: &[u64], multiplier: u64, offset: usize) -> bool { + if multiplier == 0 { + return false; + } + + let mut carry = 0_u128; + for (index, value_limb) in value.iter().copied().enumerate() { + let product = u128::from(multiplier) * u128::from(value_limb) + carry; + let product_low = product as u64; + carry = product >> 64; + + let target_index = offset + index; + let (difference, borrowed) = target[target_index].overflowing_sub(product_low); + target[target_index] = difference; + carry += u128::from(borrowed); + } + + let target_index = offset + value.len(); + let carry_low = carry as u64; + let carry_high = carry >> 64; + let (difference, borrowed) = target[target_index].overflowing_sub(carry_low); + target[target_index] = difference; + borrowed || carry_high != 0 +} + +#[inline(always)] +fn add_at(target: &mut [u64], value: &[u64], offset: usize) { + let mut carry = 0_u128; + for (index, value_limb) in value.iter().copied().enumerate() { + let target_index = offset + index; + let sum = u128::from(target[target_index]) + u128::from(value_limb) + carry; + target[target_index] = sum as u64; + carry = sum >> 64; + } + + let mut target_index = offset + value.len(); + while carry != 0 { + let sum = u128::from(target[target_index]) + carry; + target[target_index] = sum as u64; + carry = sum >> 64; + target_index += 1; + } +} + +fn shl_abs_limbs(value: &[u64], bits: usize) -> LimbVec { + let mut output = LimbVec::new(); + shl_abs_limbs_into(value, bits, &mut output); + output +} + +fn shl_abs_limbs_into(value: &[u64], bits: usize, output: &mut LimbVec) { + output.clear(); + if value.is_empty() { + return; + } + if bits == 0 { + output.extend_from_slice(value); + return; + } + + let limb_shift = bits / 64; + let bit_shift = (bits % 64) as u32; + if limb_shift == 0 && bit_shift != 0 { + shl_abs_limbs_small_into(value, bit_shift, output); + return; + } + output.resize(limb_shift + value.len() + usize::from(bit_shift != 0), 0); + let mut carry = 0_u64; + + for (index, limb) in value.iter().copied().enumerate() { + output[index + limb_shift] = if bit_shift == 0 { + limb + } else { + (limb << bit_shift) | carry + }; + carry = if bit_shift == 0 { + 0 + } else { + limb >> (64 - bit_shift) + }; + } + if bit_shift != 0 { + output[limb_shift + value.len()] = carry; + } + + trim_leading_zero_limbs(output); +} + +#[allow(clippy::uninit_vec)] +fn shl_abs_limbs_small_into(value: &[u64], bit_shift: u32, output: &mut LimbVec) { + debug_assert!((1..64).contains(&bit_shift)); + output.reserve(value.len() + 1); + // SAFETY: u64 has no drop glue, reserve guarantees room for value.len() + 1 + // limbs, and the loop initializes every slot before the vector is read. + unsafe { + output.set_len(value.len() + 1); + } + let mut carry = 0_u64; + for (index, limb) in value.iter().copied().enumerate() { + output[index] = (limb << bit_shift) | carry; + carry = limb >> (64 - bit_shift); + } + if carry != 0 { + output[value.len()] = carry; + } else { + output.truncate(value.len()); + } +} + +fn shl_abs_limbs_assign(value: &mut LimbVec, bits: usize) { + if value.is_empty() || bits == 0 { + return; + } + + let limb_shift = bits / 64; + let bit_shift = (bits % 64) as u32; + if limb_shift == 0 && bit_shift != 0 { + shl_abs_limbs_small_assign(value, bit_shift); + return; + } + let original_len = value.len(); + value.resize(original_len + limb_shift + usize::from(bit_shift != 0), 0); + + for index in (0..original_len).rev() { + let limb = value[index]; + let output_index = index + limb_shift; + value[output_index] = limb << bit_shift; + if bit_shift != 0 { + value[output_index + 1] |= limb >> (64 - bit_shift); + } + } + + value[..limb_shift].fill(0); + trim_leading_zero_limbs(value); +} + +fn shl_abs_limbs_small_assign(value: &mut LimbVec, bit_shift: u32) { + debug_assert!((1..64).contains(&bit_shift)); + let mut carry = 0_u64; + for limb in value.iter_mut() { + let current = *limb; + *limb = (current << bit_shift) | carry; + carry = current >> (64 - bit_shift); + } + if carry != 0 { + value.push(carry); + } +} + +fn shr_abs_limbs(value: &[u64], bits: usize) -> LimbVec { + let mut output = LimbVec::new(); + shr_abs_limbs_into(value, bits, &mut output); + output +} + +fn shr_abs_limbs_into(value: &[u64], bits: usize, output: &mut LimbVec) { + output.clear(); + let limb_shift = bits / 64; + if limb_shift >= value.len() { + return; + } + + let bit_shift = (bits % 64) as u32; + if limb_shift == 0 && bit_shift != 0 { + shr_abs_limbs_small_into(value, bit_shift, output); + return; + } + output.reserve(value.len() - limb_shift); + for index in limb_shift..value.len() { + let mut limb = value[index] >> bit_shift; + if bit_shift != 0 { + limb |= value.get(index + 1).copied().unwrap_or(0) << (64 - bit_shift); + } + output.push(limb); + } + + trim_leading_zero_limbs(output); +} + +fn shr_abs_limbs_small_into(value: &[u64], bit_shift: u32, output: &mut LimbVec) { + debug_assert!((1..64).contains(&bit_shift)); + output.reserve(value.len()); + for index in 0..value.len() { + let limb = (value[index] >> bit_shift) + | (value.get(index + 1).copied().unwrap_or(0) << (64 - bit_shift)); + output.push(limb); + } + trim_leading_zero_limbs(output); +} + +fn trim_leading_zero_limbs(limbs: &mut LimbVec) { + while limbs.last() == Some(&0) { + limbs.pop(); + } +} + +fn abs_bit_len(limbs: &[u64]) -> usize { + let Some(last) = limbs.last() else { + return 0; + }; + ((limbs.len() - 1) * 64) + (u64::BITS - last.leading_zeros()) as usize +} + +fn get_abs_bit(limbs: &[u64], bit: usize) -> bool { + let limb = bit / 64; + let offset = bit % 64; + limbs + .get(limb) + .map(|value| (value & (1_u64 << offset)) != 0) + .unwrap_or(false) +} + +fn set_abs_bit(limbs: &mut LimbVec, bit: usize) { + let limb = bit / 64; + let offset = bit % 64; + if limbs.len() <= limb { + limbs.resize(limb + 1, 0); + } + limbs[limb] |= 1_u64 << offset; +} + +#[cfg(test)] +mod tests { + use num_bigint::BigInt; + use num_integer::Integer; + use num_traits::Signed; + use proptest::prelude::*; + + use super::{LimbInt, LimbScratch}; + + fn arb_bigint() -> impl Strategy<Value = BigInt> { + proptest::collection::vec(any::<u64>(), 0..=24).prop_flat_map(|limbs| { + any::<bool>().prop_map(move |negative| { + let mut bytes = Vec::with_capacity(limbs.len() * 8); + for limb in &limbs { + bytes.extend_from_slice(&limb.to_le_bytes()); + } + let magnitude = num_bigint::BigUint::from_bytes_le(&bytes); + let sign = if magnitude == num_bigint::BigUint::from(0_u8) { + num_bigint::Sign::NoSign + } else if negative { + num_bigint::Sign::Minus + } else { + num_bigint::Sign::Plus + }; + BigInt::from_biguint(sign, magnitude) + }) + }) + } + + proptest! { + #![proptest_config(ProptestConfig { cases: 512, .. ProptestConfig::default() })] + + #[test] + fn roundtrips_bigint(value in arb_bigint()) { + prop_assert_eq!(LimbInt::from_bigint(&value).to_bigint(), value); + } + + #[test] + fn add_matches_bigint(left in arb_bigint(), right in arb_bigint()) { + let actual = LimbInt::from_bigint(&left).add(&LimbInt::from_bigint(&right)).to_bigint(); + prop_assert_eq!(actual, left + right); + } + + #[test] + fn sub_matches_bigint(left in arb_bigint(), right in arb_bigint()) { + let actual = LimbInt::from_bigint(&left).sub(&LimbInt::from_bigint(&right)).to_bigint(); + prop_assert_eq!(actual, left - right); + } + + #[test] + fn mul_matches_bigint(left in arb_bigint(), right in arb_bigint()) { + let actual = LimbInt::from_bigint(&left).mul(&LimbInt::from_bigint(&right)).to_bigint(); + prop_assert_eq!(actual, left * right); + } + + #[test] + fn square_matches_bigint(value in arb_bigint()) { + let actual = LimbInt::from_bigint(&value).square().to_bigint(); + prop_assert_eq!(actual, &value * &value); + } + + #[test] + fn shifts_match_positive_bigint(value in arb_bigint(), bits in 0_usize..512) { + let value = value.abs(); + let limbs = LimbInt::from_bigint(&value); + prop_assert_eq!(limbs.shl_bits(bits).to_bigint(), &value << bits); + prop_assert_eq!(limbs.shr_abs_bits(bits).to_bigint(), &value >> bits); + + let mut owned_shift = limbs.clone(); + owned_shift = owned_shift.shl_bits_owned(bits); + prop_assert_eq!(owned_shift.to_bigint(), &value << bits); + } + + #[test] + fn shifted_low_word_matches_num_bigint(value in arb_bigint(), shift in 0_u64..1536) { + let value = value.abs(); + let limbs = LimbInt::from_bigint(&value); + let mut digits = value.iter_u64_digits(); + let digit_index = usize::try_from(shift / 64).unwrap_or(usize::MAX); + let offset = (shift % 64) as u32; + let low = digits.nth(digit_index).unwrap_or(0); + let expected = if offset == 0 { + low + } else { + let high = digits.next().unwrap_or(0); + (low >> offset) | (high << (64 - offset)) + }; + + prop_assert_eq!(limbs.shifted_low_word(shift), expected); + } + + #[test] + fn div_rem_matches_bigint(left in arb_bigint(), right in arb_bigint().prop_filter("non-zero divisor", |value| value != &BigInt::from(0))) { + let left_limbs = LimbInt::from_bigint(&left); + let right_limbs = LimbInt::from_bigint(&right); + let (actual_q, actual_r) = left_limbs.div_rem(&right_limbs); + let mut scratch = LimbScratch::default(); + let actual_div = left_limbs.div_with_scratch(&right_limbs, &mut scratch); + let actual_floor = left_limbs.div_floor_with_scratch(&right_limbs, &mut scratch); + let (expected_q, expected_r) = left.div_rem(&right); + + prop_assert_eq!(actual_q.to_bigint(), expected_q); + prop_assert_eq!(actual_r.to_bigint(), expected_r); + prop_assert_eq!(actual_div.to_bigint(), actual_q.to_bigint()); + prop_assert_eq!(actual_floor.to_bigint(), left.div_floor(&right)); + } + + #[test] + fn extended_gcd_matches_bezout( + left in arb_bigint(), + right in arb_bigint().prop_filter("not both zero", |right| right != &BigInt::from(0)), + ) { + let left_limbs = LimbInt::from_bigint(&left); + let right_limbs = LimbInt::from_bigint(&right); + let actual = left_limbs.extended_gcd(&right_limbs); + let actual_x = actual.x.to_bigint(); + let actual_y = actual.y.to_bigint(); + let actual_gcd = actual.gcd.to_bigint(); + + prop_assert_eq!(&left * actual_x + &right * actual_y, actual_gcd.clone()); + prop_assert_eq!(actual_gcd.clone(), left.gcd(&right)); + } + + #[test] + fn left_extended_gcd_matches_modular_bezout( + left in arb_bigint(), + right in arb_bigint().prop_filter("non-zero modulus", |value| value != &BigInt::from(0)), + ) { + let left_limbs = LimbInt::from_bigint(&left); + let right_limbs = LimbInt::from_bigint(&right); + let mut scratch = LimbScratch::default(); + let actual = left_limbs.left_extended_gcd_with_scratch(&right_limbs, &mut scratch); + let actual_x = actual.x.to_bigint(); + let actual_gcd = actual.gcd.to_bigint(); + + prop_assert_eq!(actual_gcd.clone(), left.gcd(&right)); + prop_assert_eq!((&left * actual_x - actual_gcd) % &right, BigInt::from(0)); + } + + #[test] + fn mod_positive_matches_corrected_bigint(value in arb_bigint(), modulus in arb_bigint().prop_filter("positive modulus", |value| value > &BigInt::from(0))) { + let value_limbs = LimbInt::from_bigint(&value); + let modulus_limbs = LimbInt::from_bigint(&modulus); + let mut expected = value % &modulus; + if expected.is_negative() { + expected += &modulus; + } + + prop_assert_eq!(value_limbs.mod_positive(&modulus_limbs).to_bigint(), expected); + } + } + + #[test] + fn from_i128_handles_i128_min() { + assert_eq!( + LimbInt::from_i128(i128::MIN).to_bigint(), + BigInt::from(i128::MIN) + ); + } + + #[test] + fn xgcd_partial_matches_bigint_reference() { + let mut state = 0x811c_9dc5_0123_4567_u64; + + for case in 0..128 { + let mut r2 = next_positive_bigint(&mut state, 6); + let mut r1 = next_positive_bigint(&mut state, 6); + if r2 < r1 { + std::mem::swap(&mut r2, &mut r1); + } + + let threshold_seed = next_positive_bigint(&mut state, 3); + let threshold = (threshold_seed % &r1).max(BigInt::from(1)); + + let expected = crate::domain::vdf::arithmetic::xgcd_partial(&r2, &r1, &threshold); + let actual = super::xgcd_partial( + &LimbInt::from_bigint(&r2), + &LimbInt::from_bigint(&r1), + &LimbInt::from_bigint(&threshold), + ); + + assert_eq!(actual.0.to_bigint(), expected.0, "case={case} co2"); + assert_eq!(actual.1.to_bigint(), expected.1, "case={case} co1"); + assert_eq!(actual.2.to_bigint(), expected.2, "case={case} r2"); + assert_eq!(actual.3.to_bigint(), expected.3, "case={case} r1"); + } + } + + fn next_positive_bigint(state: &mut u64, max_limbs: usize) -> BigInt { + let limb_count = (next_u64(state) as usize % max_limbs) + 1; + let mut bytes = Vec::with_capacity(limb_count * 8); + for _ in 0..limb_count { + bytes.extend_from_slice(&next_u64(state).to_le_bytes()); + } + let value = num_bigint::BigUint::from_bytes_le(&bytes); + BigInt::from(value.max(num_bigint::BigUint::from(1_u8))) + } + + fn next_u64(state: &mut u64) -> u64 { + *state = state + .wrapping_mul(6_364_136_223_846_793_005) + .wrapping_add(1); + *state + } +} diff --git a/src/domain/vdf/mod.rs b/src/domain/vdf/mod.rs @@ -0,0 +1,216 @@ +use std::time::{Duration, Instant}; + +use super::{Block, FinalizerMode, MAX_VDF_ROUNDS, VDF_TARGET_BLOCK_MS, decode_hex, hex_encode}; + +#[cfg(test)] +mod arithmetic; +mod limb_arithmetic; +mod limbs; +mod prover; +#[cfg(test)] +mod reducer; +mod wesolowski; + +#[cfg(test)] +mod reference; + +const VDF_SOLUTION_PREFIX: &str = "classgroup-wesolowski-bqfc-v1:"; +const MIN_VDF_ROUNDS: u64 = 1; +pub(super) const VDF_RETARGET_WINDOW_BLOCKS: usize = 20; +pub(super) const MAX_VDF_RETARGET_STEP_PERCENT: u128 = 2; +pub(super) const VDF_RETARGET_DEADBAND_PERCENT: u128 = 10; +pub(super) const MIN_VDF_RETARGET_OBSERVED_BLOCK_MS: u64 = VDF_TARGET_BLOCK_MS / 4; +pub(super) const MAX_VDF_RETARGET_OBSERVED_BLOCK_MS: u64 = VDF_TARGET_BLOCK_MS * 4; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct VdfProgress { + pub completed_steps: u64, + pub total_steps: u64, + pub completed_phase_rounds: u64, + pub phase_rounds: u64, + pub phase: VdfProgressPhase, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub enum VdfProgressPhase { + Output, + Proof, +} + +pub fn run_vdf(seed: &str, rounds: u64) -> String { + run_vdf_with_progress(seed, rounds, Duration::MAX, |_| {}) +} + +pub fn run_vdf_with_progress( + seed: &str, + rounds: u64, + progress_interval: Duration, + mut progress: impl FnMut(VdfProgress), +) -> String { + let total_steps = rounds.saturating_mul(2); + let mut last_progress = Instant::now(); + let solution = wesolowski::prove(seed.as_bytes(), rounds, |phase, completed_phase_rounds| { + let completed_steps = match phase { + VdfProgressPhase::Output => completed_phase_rounds, + VdfProgressPhase::Proof => rounds.saturating_add(completed_phase_rounds), + }; + maybe_report_vdf_progress( + &mut last_progress, + progress_interval, + VdfProgress { + completed_steps, + total_steps, + completed_phase_rounds, + phase_rounds: rounds, + phase, + }, + &mut progress, + ); + }) + .expect("valid IUNA VDF parameters must produce a class-group proof"); + encode_vdf_solution(&solution) +} + +pub fn verify_vdf(seed: &str, rounds: u64, solution: &str) -> bool { + let Some(solution) = decode_vdf_solution(solution) else { + return false; + }; + wesolowski::verify(seed.as_bytes(), rounds, &solution) +} + +pub(super) fn vdf_solution_placeholder() -> String { + format!( + "{VDF_SOLUTION_PREFIX}{}", + "f".repeat(wesolowski::SOLUTION_BYTES * 2) + ) +} + +pub(super) fn retarget_vdf_rounds(current_rounds: u64, observed_block_ms: u64) -> u64 { + let current = u128::from(current_rounds); + let observed = u128::from(observed_block_ms.max(1)); + let target = u128::from(VDF_TARGET_BLOCK_MS); + let deadband = target * VDF_RETARGET_DEADBAND_PERCENT / 100; + if observed >= target.saturating_sub(deadband) && observed <= target.saturating_add(deadband) { + return current_rounds; + } + + let raw_adjusted = current * target / observed; + let max_step = (current * MAX_VDF_RETARGET_STEP_PERCENT / 100).max(1); + let min_next = current + .saturating_sub(max_step) + .max(u128::from(MIN_VDF_ROUNDS)); + let max_next = current + .saturating_add(max_step) + .min(u128::from(MAX_VDF_ROUNDS)); + raw_adjusted.clamp(min_next, max_next) as u64 +} + +pub(super) fn clamped_vdf_retarget_observed_block_ms(observed_block_ms: u64) -> u64 { + observed_block_ms.clamp( + MIN_VDF_RETARGET_OBSERVED_BLOCK_MS, + MAX_VDF_RETARGET_OBSERVED_BLOCK_MS, + ) +} + +pub(super) fn vdf_retarget_observed_block_ms(parent: &Block, child: &Block) -> Option<u64> { + if child.finalizer_mode != FinalizerMode::Ticket { + return None; + } + if child.finalizer_rank != 0 { + return None; + } + + Some(clamped_vdf_retarget_observed_block_ms( + child.timestamp_ms - parent.timestamp_ms, + )) +} + +fn maybe_report_vdf_progress( + last_progress: &mut Instant, + progress_interval: Duration, + snapshot: VdfProgress, + progress: &mut impl FnMut(VdfProgress), +) { + if snapshot.completed_steps == snapshot.total_steps + || last_progress.elapsed() >= progress_interval + { + progress(snapshot); + *last_progress = Instant::now(); + } +} + +fn encode_vdf_solution(solution: &[u8]) -> String { + format!("{VDF_SOLUTION_PREFIX}{}", hex_encode(solution)) +} + +fn decode_vdf_solution(solution: &str) -> Option<Vec<u8>> { + let hex = solution.strip_prefix(VDF_SOLUTION_PREFIX)?; + let solution = decode_hex(hex).ok()?; + (solution.len() == wesolowski::SOLUTION_BYTES).then_some(solution) +} + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use super::{ + VDF_SOLUTION_PREFIX, VdfProgressPhase, run_vdf, run_vdf_with_progress, + vdf_solution_placeholder, verify_vdf, wesolowski, + }; + + #[test] + fn vdf_solution_verifies_and_is_bound_to_seed_and_rounds() { + let solution = run_vdf("test-seed", 128); + + assert!(verify_vdf("test-seed", 128, &solution)); + assert!(!verify_vdf("other-seed", 128, &solution)); + assert!(!verify_vdf("test-seed", 129, &solution)); + assert!(!verify_vdf("test-seed", 128, "not-a-vdf-solution")); + } + + #[test] + fn vdf_progress_reports_output_and_proof_steps() { + let mut progress = Vec::new(); + let solution = run_vdf_with_progress("progress-seed", 4, Duration::ZERO, |snapshot| { + progress.push(snapshot); + }); + + assert!(verify_vdf("progress-seed", 4, &solution)); + assert!( + progress + .iter() + .any(|snapshot| snapshot.phase == VdfProgressPhase::Output) + ); + assert!( + progress + .iter() + .any(|snapshot| snapshot.phase == VdfProgressPhase::Proof) + ); + assert_eq!( + progress.last().map(|snapshot| snapshot.completed_steps), + Some(8) + ); + assert_eq!( + progress.last().map(|snapshot| snapshot.total_steps), + Some(8) + ); + } + + #[test] + fn vdf_solution_uses_chia_bqfc_protocol_format() { + let solution = run_vdf("test-seed", 16); + let encoded = solution.strip_prefix(VDF_SOLUTION_PREFIX).unwrap(); + + assert_eq!(encoded.len(), wesolowski::SOLUTION_BYTES * 2); + assert_eq!(solution.len(), vdf_solution_placeholder().len()); + } + + #[test] + fn legacy_vdf_solution_formats_are_not_accepted() { + let rsa = format!("{}:{}", "f".repeat(512), "f".repeat(512)); + let gmp_class_group = format!("classgroup-wesolowski-v1:{}", "f".repeat(520)); + + assert!(!verify_vdf("test-seed", 16, &rsa)); + assert!(!verify_vdf("test-seed", 16, &gmp_class_group)); + } +} diff --git a/src/domain/vdf/prover.rs b/src/domain/vdf/prover.rs @@ -0,0 +1,591 @@ +use kyn_vdf::{Form, KynVdfError, get_b}; +use num_bigint::{BigInt, BigUint}; +use num_traits::{One, ToPrimitive}; + +use super::{VdfProgressPhase, limb_arithmetic}; + +// Keep peak prover memory bounded. Larger parameter sets retain the old +// constant-memory algorithm instead of attempting an attacker-sized allocation. +const MAX_CHECKPOINTS: u64 = 262_144; +const MAX_BUCKETS: u64 = 65_536; +const INVALID_BUCKET: usize = usize::MAX; + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +struct ProofParameters { + k: u32, + l: u64, + checkpoint_count: u64, + bucket_count: u64, +} + +#[derive(Clone, Copy)] +struct ClassGroup<'a> { + discriminant: &'a BigInt, + threshold: &'a BigInt, +} + +pub(super) fn prove( + discriminant: &BigInt, + generator: &Form, + threshold: &BigInt, + rounds: u64, + progress: impl FnMut(VdfProgressPhase, u64), +) -> Result<(Form, Form), KynVdfError> { + let group = ClassGroup { + discriminant, + threshold, + }; + let parameters = ProofParameters::for_rounds(rounds); + if parameters.checkpoint_count > MAX_CHECKPOINTS || parameters.bucket_count > MAX_BUCKETS { + return prove_constant_memory(discriminant, generator, threshold, rounds, progress); + } + + prove_checkpointed(group, generator, rounds, parameters, progress) +} + +fn prove_checkpointed( + group: ClassGroup<'_>, + generator: &Form, + rounds: u64, + parameters: ProofParameters, + mut progress: impl FnMut(VdfProgressPhase, u64), +) -> Result<(Form, Form), KynVdfError> { + let checkpoint_capacity = usize::try_from(parameters.checkpoint_count) + .map_err(|_| arithmetic_error("checkpoint count does not fit in memory"))?; + let checkpoint_stride = u64::from(parameters.k) + .checked_mul(parameters.l) + .ok_or_else(|| arithmetic_error("checkpoint stride overflow"))?; + let limb_discriminant = limb_arithmetic::to_limb(group.discriminant); + let limb_threshold = limb_arithmetic::to_limb(group.threshold); + let mut checkpoints = Vec::with_capacity(checkpoint_capacity); + let mut output = limb_arithmetic::LimbForm::from_form(generator); + let mut output_scratch = limb_arithmetic::LimbFormScratch::default(); + for completed_rounds in 1..=rounds { + if (completed_rounds - 1) % checkpoint_stride == 0 { + checkpoints.push(output.clone()); + } + output = output.nudupl_reduce_with_scratch( + &limb_discriminant, + &limb_threshold, + &mut output_scratch, + ); + progress(VdfProgressPhase::Output, completed_rounds); + } + + let output = output.into_form(); + debug_assert_eq!(checkpoints.len(), checkpoint_capacity); + let proof = generate_checkpoint_proof( + group, + generator, + &output, + &checkpoints, + rounds, + parameters, + &mut progress, + )?; + Ok((output, proof)) +} + +fn generate_checkpoint_proof( + group: ClassGroup<'_>, + generator: &Form, + output: &Form, + checkpoints: &[limb_arithmetic::LimbForm], + rounds: u64, + parameters: ProofParameters, + progress: &mut impl FnMut(VdfProgressPhase, u64), +) -> Result<Form, KynVdfError> { + let challenge = get_b(group.discriminant, generator, output)?; + let bucket_count = usize::try_from(parameters.bucket_count) + .map_err(|_| arithmetic_error("bucket count does not fit in memory"))?; + let k0 = parameters.k - parameters.k / 2; + let k1 = parameters.k / 2; + let row_count = 1_u64 << k1; + let column_count = 1_u64 << k0; + let work_per_pass = parameters + .checkpoint_count + .checked_add(row_count) + .and_then(|value| value.checked_add(column_count)) + .ok_or_else(|| arithmetic_error("proof progress calculation overflow"))?; + let total_work = parameters + .l + .checked_mul(work_per_pass) + .ok_or_else(|| arithmetic_error("proof progress calculation overflow"))?; + let mut completed_work = 0_u64; + let limb_discriminant = limb_arithmetic::to_limb(group.discriminant); + let limb_threshold = limb_arithmetic::to_limb(group.threshold); + let mut proof = limb_arithmetic::LimbForm::identity(&limb_discriminant); + let mut proof_scratch = limb_arithmetic::LimbFormScratch::default(); + let block_step_exponent = u64::from(parameters.k) + .checked_mul(parameters.l) + .ok_or_else(|| arithmetic_error("proof block step overflow"))?; + let block_step = BigUint::from(2_u8).modpow(&BigUint::from(block_step_exponent), &challenge); + + for j in (0..parameters.l).rev() { + proof = proof.fast_pow_u64_with_scratch( + 1_u64 << parameters.k, + &limb_discriminant, + &limb_threshold, + &mut proof_scratch, + ); + + let checkpoint_blocks = + get_blocks_for_pass(j, parameters, rounds, &challenge, &block_step)?; + let mut buckets: Vec<Option<limb_arithmetic::LimbForm>> = vec![None; bucket_count]; + for (i, checkpoint) in checkpoints.iter().enumerate() { + let bucket = checkpoint_blocks[i]; + if bucket != INVALID_BUCKET { + let bucket_form = buckets[bucket].take(); + buckets[bucket] = Some(match bucket_form { + Some(bucket_form) => bucket_form.compose_unreduced( + checkpoint, + &limb_discriminant, + &limb_threshold, + &mut proof_scratch, + ), + None => checkpoint.clone(), + }); + } + + completed_work += 1; + report_proof_progress(progress, rounds, completed_work, total_work); + } + + for b1 in 0..row_count { + let row_start = b1 << k0; + let mut aggregate: Option<limb_arithmetic::LimbForm> = None; + for b0 in 0..column_count { + if let Some(bucket) = &buckets[(row_start + b0) as usize] { + aggregate = Some(match aggregate { + Some(aggregate) => aggregate.compose_unreduced( + bucket, + &limb_discriminant, + &limb_threshold, + &mut proof_scratch, + ), + None => bucket.clone(), + }); + } + } + if row_start != 0 { + if let Some(aggregate) = aggregate { + let aggregate = aggregate.fast_pow_u64_with_scratch( + row_start, + &limb_discriminant, + &limb_threshold, + &mut proof_scratch, + ); + proof = proof.compose_unreduced( + &aggregate, + &limb_discriminant, + &limb_threshold, + &mut proof_scratch, + ); + } + } + completed_work += 1; + report_proof_progress(progress, rounds, completed_work, total_work); + } + + for b0 in 0..column_count { + let mut aggregate: Option<limb_arithmetic::LimbForm> = None; + for b1 in 0..row_count { + if let Some(bucket) = &buckets[((b1 << k0) + b0) as usize] { + aggregate = Some(match aggregate { + Some(aggregate) => aggregate.compose_unreduced( + bucket, + &limb_discriminant, + &limb_threshold, + &mut proof_scratch, + ), + None => bucket.clone(), + }); + } + } + if b0 != 0 { + if let Some(aggregate) = aggregate { + let aggregate = aggregate.fast_pow_u64_with_scratch( + b0, + &limb_discriminant, + &limb_threshold, + &mut proof_scratch, + ); + proof = proof.compose_unreduced( + &aggregate, + &limb_discriminant, + &limb_threshold, + &mut proof_scratch, + ); + } + } + completed_work += 1; + report_proof_progress(progress, rounds, completed_work, total_work); + } + } + + proof.reduce(); + progress(VdfProgressPhase::Proof, rounds); + Ok(proof.into_form()) +} + +fn get_blocks_for_pass( + j: u64, + parameters: ProofParameters, + rounds: u64, + challenge: &BigUint, + step: &BigUint, +) -> Result<Vec<usize>, KynVdfError> { + let checkpoint_count = usize::try_from(parameters.checkpoint_count) + .map_err(|_| arithmetic_error("checkpoint count does not fit in memory"))?; + let mut blocks = vec![INVALID_BUCKET; checkpoint_count]; + if checkpoint_count == 0 { + return Ok(blocks); + } + + let mut index = checkpoint_count - 1; + loop { + let position = checkpoint_position(index, j, parameters.l)?; + if rounds >= block_end(position, parameters.k)? { + break; + } + if index == 0 { + return Ok(blocks); + } + index -= 1; + } + + let position = checkpoint_position(index, j, parameters.l)?; + let exponent = rounds + .checked_sub(block_end(position, parameters.k)?) + .ok_or_else(|| arithmetic_error("proof block exceeds iteration count"))?; + let mut residue = BigUint::from(2_u8).modpow(&BigUint::from(exponent), challenge); + loop { + blocks[index] = block_from_residue(&residue, parameters.k, challenge)?; + if index == 0 { + break; + } + index -= 1; + residue = (residue * step) % challenge; + } + + Ok(blocks) +} + +fn report_proof_progress( + progress: &mut impl FnMut(VdfProgressPhase, u64), + rounds: u64, + completed: u64, + total: u64, +) { + let proof_rounds = completed.saturating_mul(rounds) / total.max(1); + progress(VdfProgressPhase::Proof, proof_rounds); +} + +fn checkpoint_position(index: usize, j: u64, l: u64) -> Result<u64, KynVdfError> { + u64::try_from(index) + .map_err(|_| arithmetic_error("checkpoint index does not fit in u64"))? + .checked_mul(l) + .and_then(|value| value.checked_add(j)) + .ok_or_else(|| arithmetic_error("checkpoint position overflow")) +} + +fn block_end(position: u64, k: u32) -> Result<u64, KynVdfError> { + u64::from(k) + .checked_mul( + position + .checked_add(1) + .ok_or_else(|| arithmetic_error("checkpoint position overflow"))?, + ) + .ok_or_else(|| arithmetic_error("checkpoint block overflow")) +} + +fn block_from_residue( + residue: &BigUint, + k: u32, + challenge: &BigUint, +) -> Result<usize, KynVdfError> { + let block = (residue << k) / challenge; + block + .to_usize() + .filter(|value| *value < (1_usize << k)) + .ok_or_else(|| arithmetic_error("proof block does not fit in its bucket range")) +} + +fn prove_constant_memory( + discriminant: &BigInt, + generator: &Form, + threshold: &BigInt, + rounds: u64, + mut progress: impl FnMut(VdfProgressPhase, u64), +) -> Result<(Form, Form), KynVdfError> { + let limb_discriminant = limb_arithmetic::to_limb(discriminant); + let limb_threshold = limb_arithmetic::to_limb(threshold); + let mut output = limb_arithmetic::LimbForm::from_form(generator); + let mut output_scratch = limb_arithmetic::LimbFormScratch::default(); + for completed_rounds in 1..=rounds { + output = output.nudupl_reduce_with_scratch( + &limb_discriminant, + &limb_threshold, + &mut output_scratch, + ); + progress(VdfProgressPhase::Output, completed_rounds); + } + let output = output.into_form(); + + let challenge = get_b(discriminant, generator, &output)?; + let generator_limb = limb_arithmetic::LimbForm::from_form(generator); + let mut proof = limb_arithmetic::LimbForm::identity(&limb_discriminant); + let mut proof_scratch = limb_arithmetic::LimbFormScratch::default(); + let mut remainder = BigUint::one() % &challenge; + for completed_rounds in 1..=rounds { + let doubled = &remainder << 1_usize; + let carry = doubled >= challenge; + proof = proof.nudupl_reduce_with_scratch( + &limb_discriminant, + &limb_threshold, + &mut proof_scratch, + ); + if carry { + proof = proof.nucomp_reduce_with_scratch( + &generator_limb, + &limb_discriminant, + &limb_threshold, + &mut proof_scratch, + ); + } + remainder = doubled % &challenge; + progress(VdfProgressPhase::Proof, completed_rounds); + } + + Ok((output, proof.into_form())) +} + +fn arithmetic_error(message: &str) -> KynVdfError { + KynVdfError::ArithmeticError(message.to_owned()) +} + +impl ProofParameters { + #[allow(clippy::approx_constant)] // Match Chia's published parameter heuristic. + fn for_rounds(rounds: u64) -> Self { + let log_memory = 23.253_496_66_f64; + let log_rounds = (rounds as f64).log2(); + let l = if log_rounds - log_memory > 0.000_001 { + 2_f64.powf(log_memory - 20.0).ceil() as u64 + } else { + 1 + }; + let intermediate = rounds as f64 * 0.693_147_1 / (2.0 * l as f64); + let mut k = if intermediate <= 1.0 { + 1 + } else { + (intermediate.ln() - intermediate.ln().ln() + 0.25) + .round() + .max(1.0) as u32 + }; + if rounds >= 100_000 { + k = k.max(10); + } + let checkpoint_stride = u64::from(k).saturating_mul(l).max(1); + + Self { + k, + l, + checkpoint_count: rounds.div_ceil(checkpoint_stride), + bucket_count: 1_u64.checked_shl(k).unwrap_or(u64::MAX), + } + } +} + +#[cfg(test)] +mod tests { + use std::time::Instant; + + use kyn_vdf::{Form, create_discriminant, isqrt_fourth}; + use num_traits::Signed; + + use super::{ + ClassGroup, ProofParameters, VdfProgressPhase, prove, prove_checkpointed, + prove_constant_memory, + }; + + #[test] + fn parameters_match_chia_reference_values() { + for (rounds, k, l, checkpoints, buckets) in [ + (1, 1, 1, 1, 2), + (100, 3, 1, 34, 8), + (300, 3, 1, 100, 8), + (100_000, 10, 1, 10_000, 1_024), + (1_000_000, 10, 1, 100_000, 1_024), + (1_500_000, 11, 1, 136_364, 2_048), + (20_000_000, 11, 10, 181_819, 2_048), + ] { + assert_eq!( + ProofParameters::for_rounds(rounds), + ProofParameters { + k, + l, + checkpoint_count: checkpoints, + bucket_count: buckets, + } + ); + } + } + + #[test] + fn checkpoint_proof_matches_constant_memory_proof() { + let discriminant = create_discriminant(b"iuna-vdf-checkpoint-differential", 1024).unwrap(); + let generator = Form::generator(&discriminant).unwrap(); + let threshold = isqrt_fourth(&discriminant.abs()); + + for rounds in [1, 2, 16, 100, 300, 1_001] { + let checkpoint = + prove(&discriminant, &generator, &threshold, rounds, |_, _| {}).unwrap(); + let constant_memory = + prove_constant_memory(&discriminant, &generator, &threshold, rounds, |_, _| {}) + .unwrap(); + + assert_eq!(checkpoint, constant_memory, "rounds={rounds}"); + } + } + + #[test] + fn multi_pass_checkpoint_proof_matches_constant_memory_and_reports_monotonic_progress() { + let rounds = 301_u64; + let discriminant = create_discriminant(b"iuna-vdf-checkpoint-multi-pass", 1024).unwrap(); + let generator = Form::generator(&discriminant).unwrap(); + let threshold = isqrt_fourth(&discriminant.abs()); + let group = ClassGroup { + discriminant: &discriminant, + threshold: &threshold, + }; + let parameters = ProofParameters { + k: 3, + l: 2, + checkpoint_count: rounds.div_ceil(6), + bucket_count: 8, + }; + let mut proof_progress = Vec::new(); + + let checkpoint = + prove_checkpointed(group, &generator, rounds, parameters, |phase, completed| { + if phase == VdfProgressPhase::Proof { + proof_progress.push(completed); + } + }) + .unwrap(); + let constant_memory = + prove_constant_memory(&discriminant, &generator, &threshold, rounds, |_, _| {}) + .unwrap(); + + assert_eq!(checkpoint, constant_memory); + assert!(proof_progress.windows(2).all(|pair| pair[0] <= pair[1])); + assert_eq!(proof_progress.last(), Some(&rounds)); + } + + #[test] + #[ignore = "manual VDF prover benchmark"] + fn benchmark_checkpoint_prover_against_constant_memory() { + let rounds = 100_000_u64; + let discriminant = create_discriminant(b"iuna-vdf-prover-benchmark", 1024).unwrap(); + let generator = Form::generator(&discriminant).unwrap(); + let threshold = isqrt_fourth(&discriminant.abs()); + + let started = Instant::now(); + let checkpoint = prove(&discriminant, &generator, &threshold, rounds, |_, _| {}).unwrap(); + let checkpoint_elapsed = started.elapsed(); + let started = Instant::now(); + let constant_memory = + prove_constant_memory(&discriminant, &generator, &threshold, rounds, |_, _| {}) + .unwrap(); + let constant_memory_elapsed = started.elapsed(); + + assert_eq!(checkpoint, constant_memory); + eprintln!( + "rounds={rounds} checkpoint={checkpoint_elapsed:?} constant_memory={constant_memory_elapsed:?} speedup={:.2}x", + constant_memory_elapsed.as_secs_f64() / checkpoint_elapsed.as_secs_f64() + ); + } + + #[test] + #[ignore = "manual VDF phase benchmark"] + fn benchmark_checkpoint_prover_phases() { + let rounds = 100_000_u64; + let discriminant = create_discriminant(b"iuna-vdf-prover-benchmark", 1024).unwrap(); + let generator = Form::generator(&discriminant).unwrap(); + let threshold = isqrt_fourth(&discriminant.abs()); + let group = ClassGroup { + discriminant: &discriminant, + threshold: &threshold, + }; + let parameters = ProofParameters::for_rounds(rounds); + let checkpoint_stride = u64::from(parameters.k) * parameters.l; + let mut checkpoints = Vec::with_capacity(parameters.checkpoint_count as usize); + let limb_discriminant = crate::domain::vdf::limb_arithmetic::to_limb(&discriminant); + let limb_threshold = crate::domain::vdf::limb_arithmetic::to_limb(&threshold); + let mut output = crate::domain::vdf::limb_arithmetic::LimbForm::from_form(&generator); + let mut output_scratch = crate::domain::vdf::limb_arithmetic::LimbFormScratch::default(); + + let started = Instant::now(); + for completed_rounds in 1..=rounds { + if (completed_rounds - 1) % checkpoint_stride == 0 { + checkpoints.push(output.clone()); + } + output = output.nudupl_reduce_with_scratch( + &limb_discriminant, + &limb_threshold, + &mut output_scratch, + ); + } + let output = output.into_form(); + let output_elapsed = started.elapsed(); + + let started = Instant::now(); + let proof = super::generate_checkpoint_proof( + group, + &generator, + &output, + &checkpoints, + rounds, + parameters, + &mut |_, _| {}, + ) + .unwrap(); + let proof_elapsed = started.elapsed(); + + assert_eq!( + super::prove(&discriminant, &generator, &threshold, rounds, |_, _| {}).unwrap(), + (output, proof) + ); + eprintln!( + "rounds={rounds} output={output_elapsed:?} proof={proof_elapsed:?} total={:?}", + output_elapsed + proof_elapsed + ); + } + #[test] + #[ignore = "manual VDF square phase benchmark"] + fn benchmark_output_square_phases() { + let rounds = 100_000; + let discriminant = create_discriminant(b"iuna-vdf-prover-benchmark", 1024).unwrap(); + let threshold = isqrt_fourth(&discriminant.abs()); + let mut output = Form::generator(&discriminant).unwrap(); + let mut nudupl_elapsed = std::time::Duration::ZERO; + let mut reduce_elapsed = std::time::Duration::ZERO; + + for _ in 0..rounds { + let started = Instant::now(); + output = + crate::domain::vdf::arithmetic::nudupl_owned(output, &discriminant, &threshold); + nudupl_elapsed += started.elapsed(); + + let started = Instant::now(); + crate::domain::vdf::reducer::reduce(&mut output); + reduce_elapsed += started.elapsed(); + } + + assert!(output.is_reduced()); + eprintln!( + "rounds={rounds} nudupl={nudupl_elapsed:?} reduce={reduce_elapsed:?} total={:?}", + nudupl_elapsed + reduce_elapsed + ); + } +} diff --git a/src/domain/vdf/reducer.rs b/src/domain/vdf/reducer.rs @@ -0,0 +1,186 @@ +use std::mem; + +use kyn_vdf::Form; +use num_bigint::BigInt; +use num_integer::Integer; +use num_traits::{Signed, Zero}; + +const APPROXIMATION_EXPONENT_SPREAD: u64 = 31; +const TRANSFORM_COEFFICIENT_LIMIT: i128 = 1_i128 << 31; + +pub(super) fn reduce(form: &mut Form) { + while !finish_if_reduced(form) { + let (a, a_exponent) = signed_63_bit_approximation(&form.a); + let (b, b_exponent) = signed_63_bit_approximation(&form.b); + let (c, c_exponent) = signed_63_bit_approximation(&form.c); + let min_exponent = a_exponent.min(b_exponent).min(c_exponent); + let max_exponent = a_exponent.max(b_exponent).max(c_exponent); + + if max_exponent - min_exponent > APPROXIMATION_EXPONENT_SPREAD { + reduce_once(form); + continue; + } + + let common_exponent = max_exponent + 1; + let a = signed_shift(a, a_exponent as i64 - common_exponent as i64); + let b = signed_shift(b, b_exponent as i64 - common_exponent as i64); + let c = signed_shift(c, c_exponent as i64 - common_exponent as i64); + let transform = approximate_transform(a, b, c); + apply_transform(form, transform); + } +} + +fn finish_if_reduced(form: &mut Form) -> bool { + if form.a.abs() < form.b.abs() || form.c.abs() < form.b.abs() { + return false; + } + + if form.a > form.c { + mem::swap(&mut form.a, &mut form.c); + form.b = -mem::take(&mut form.b); + } else if form.a == form.c && form.b.is_negative() { + form.b = -mem::take(&mut form.b); + } + true +} + +fn reduce_once(form: &mut Form) { + let two_c = &form.c << 1_usize; + let s = (&form.b + &form.c).div_floor(&two_c); + let old_a = mem::take(&mut form.a); + let old_b = mem::take(&mut form.b); + let old_c = mem::take(&mut form.c); + let c_times_s = &old_c * &s; + + form.a = old_c; + form.b = (&c_times_s << 1_usize) - &old_b; + form.c = old_a + &s * (c_times_s - old_b); +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +struct Transform { + u: i128, + v: i128, + w: i128, + x: i128, +} + +fn approximate_transform(mut a: i128, mut b: i128, mut c: i128) -> Transform { + let mut current = Transform { + u: 1, + v: 0, + w: 0, + x: 1, + }; + + loop { + let s = if b >= 0 { + (b + c) / (c << 1) + } else { + -((-b + c) / (c << 1)) + }; + let old_a = a; + let old_b = b; + a = c; + b = -b + ((c * s) << 1); + c = old_a - s * (old_b - c * s); + + let next = Transform { + u: current.v, + v: -current.u + s * current.v, + w: current.x, + x: -current.w + s * current.x, + }; + let coefficients_fit = (next.v.abs() | next.x.abs()) <= TRANSFORM_COEFFICIENT_LIMIT; + if coefficients_fit { + current = next; + } + if !coefficients_fit || a <= c || c <= 0 { + return current; + } + } +} + +fn apply_transform(form: &mut Form, transform: Transform) { + let old_a = mem::take(&mut form.a); + let old_b = mem::take(&mut form.b); + let old_c = mem::take(&mut form.c); + let Transform { u, v, w, x } = transform; + + form.a = scaled(&old_a, u * u) + scaled(&old_b, u * w) + scaled(&old_c, w * w); + form.b = scaled(&old_a, 2 * u * v) + scaled(&old_b, u * x + v * w) + scaled(&old_c, 2 * w * x); + form.c = scaled(&old_a, v * v) + scaled(&old_b, v * x) + scaled(&old_c, x * x); +} + +fn scaled(value: &BigInt, scalar: i128) -> BigInt { + value * BigInt::from(scalar) +} + +fn signed_63_bit_approximation(value: &BigInt) -> (i128, u64) { + if value.is_zero() { + return (0, 0); + } + + let mut digits = value.iter_u64_digits(); + let digit_count = digits.len() as u64; + let top = digits + .next_back() + .expect("a nonzero BigInt has a top digit"); + let top_bits = u64::from(64 - top.leading_zeros()); + let exponent = top_bits + (digit_count - 1) * 64; + let mut approximation = if top_bits == 64 { + top >> 1 + } else { + top << (63 - top_bits) + }; + if let Some(previous) = digits.next_back() { + let shift = top_bits + 1; + if shift < 64 { + approximation += previous >> shift; + } + } + + let approximation = i128::from(approximation); + if value.is_negative() { + (-approximation, exponent) + } else { + (approximation, exponent) + } +} + +fn signed_shift(value: i128, shift: i64) -> i128 { + if shift > 0 { + value << shift + } else if shift <= -128 { + 0 + } else { + value >> -shift + } +} + +#[cfg(test)] +mod tests { + use kyn_vdf::{Form, create_discriminant, isqrt_fourth}; + use num_traits::Signed; + + use super::reduce; + + #[test] + fn pulmark_reducer_matches_canonical_reducer_across_sequential_squares() { + let discriminant = create_discriminant(b"iuna-pulmark-reducer-differential", 1024).unwrap(); + let threshold = isqrt_fourth(&discriminant.abs()); + let mut form = Form::generator(&discriminant).unwrap(); + + for round in 1..=10_000 { + let unreduced = form.nudupl(&discriminant, &threshold); + let mut expected = unreduced.clone(); + expected.reduce(&discriminant); + let mut actual = unreduced; + reduce(&mut actual); + + assert_eq!(actual, expected, "round={round}"); + assert!(actual.is_reduced(), "round={round}"); + form = actual; + } + } +} diff --git a/src/domain/vdf/reference.rs b/src/domain/vdf/reference.rs @@ -0,0 +1,37 @@ +use kyn_vdf::{Form, KynVdfError, create_discriminant, get_b}; +use num_bigint::BigUint; +use num_traits::One; + +use super::wesolowski::{DISCRIMINANT_BITS, serialize_solution}; + +fn prove(seed: &[u8], rounds: u64) -> Result<Vec<u8>, KynVdfError> { + if rounds == 0 { + return Err(KynVdfError::InvalidIterations(rounds)); + } + let shift = usize::try_from(rounds).map_err(|_| KynVdfError::InvalidIterations(rounds))?; + let discriminant = create_discriminant(seed, DISCRIMINANT_BITS)?; + let generator = + Form::generator(&discriminant).ok_or(KynVdfError::InvalidDiscriminantIdentity)?; + let exponent = BigUint::one() << shift; + let output = generator.pow(&exponent, &discriminant); + let challenge = get_b(&discriminant, &generator, &output)?; + let quotient = &exponent / challenge; + let proof = generator.pow(&quotient, &discriminant); + serialize_solution(&output, &proof) +} + +#[cfg(test)] +mod tests { + use super::prove as reference_prove; + use crate::domain::vdf::wesolowski::prove; + + #[test] + fn sequential_prover_matches_independent_exponent_oracle() { + for rounds in [1, 2, 16, 100, 300] { + let sequential = prove(b"iuna-vdf-reference-oracle", rounds, |_, _| {}).unwrap(); + let reference = reference_prove(b"iuna-vdf-reference-oracle", rounds).unwrap(); + + assert_eq!(sequential, reference, "rounds={rounds}"); + } + } +} diff --git a/src/domain/vdf/wesolowski.rs b/src/domain/vdf/wesolowski.rs @@ -0,0 +1,130 @@ +use kyn_vdf::{ + Form, KynVdfError, create_discriminant, deserialize_form, isqrt_fourth, serialize_form, + verify_wesolowski, +}; +use num_traits::Signed; + +use super::{VdfProgressPhase, prover}; + +pub(super) const DISCRIMINANT_BITS: usize = 1024; +const FORM_BYTES: usize = 100; +pub(super) const SOLUTION_BYTES: usize = FORM_BYTES * 2; + +pub(super) fn prove( + seed: &[u8], + rounds: u64, + mut progress: impl FnMut(VdfProgressPhase, u64), +) -> Result<Vec<u8>, KynVdfError> { + if rounds == 0 { + return Err(KynVdfError::InvalidIterations(rounds)); + } + + let discriminant = create_discriminant(seed, DISCRIMINANT_BITS)?; + let generator = + Form::generator(&discriminant).ok_or(KynVdfError::InvalidDiscriminantIdentity)?; + let threshold = isqrt_fourth(&discriminant.abs()); + + let (output, proof) = + prover::prove(&discriminant, &generator, &threshold, rounds, &mut progress)?; + serialize_solution(&output, &proof) +} + +pub(super) fn verify(seed: &[u8], rounds: u64, solution: &[u8]) -> bool { + if rounds == 0 || solution.len() != SOLUTION_BYTES { + return false; + } + + let Some(discriminant) = create_discriminant(seed, DISCRIMINANT_BITS).ok() else { + return false; + }; + let Some(generator) = Form::generator(&discriminant) else { + return false; + }; + + let (output_bytes, proof_bytes) = solution.split_at(FORM_BYTES); + let Some(output) = canonical_form(&discriminant, output_bytes) else { + return false; + }; + let Some(proof) = canonical_form(&discriminant, proof_bytes) else { + return false; + }; + + verify_wesolowski(&discriminant, &generator, &output, &proof, rounds).unwrap_or(false) +} + +fn canonical_form(discriminant: &num_bigint::BigInt, bytes: &[u8]) -> Option<Form> { + let form = deserialize_form(discriminant, bytes).ok()?; + (serialize_form(&form, DISCRIMINANT_BITS).ok()?.as_slice() == bytes).then_some(form) +} + +pub(super) fn serialize_solution(output: &Form, proof: &Form) -> Result<Vec<u8>, KynVdfError> { + let mut solution = serialize_form(output, DISCRIMINANT_BITS)?; + solution.extend_from_slice(&serialize_form(proof, DISCRIMINANT_BITS)?); + debug_assert_eq!(solution.len(), SOLUTION_BYTES); + Ok(solution) +} + +#[cfg(test)] +mod tests { + use super::{SOLUTION_BYTES, prove, verify}; + + const CHIA_CHALLENGE_42_PROOF_HEX: &str = concat!( + "0300032167dfd0eb393ed5d544e6499ba24def860ecd8a3600490f2f87b003c3e7855763969d34e2d1c60910297df3aead9f078a1f4d3973903f532977f9639f693cdbd331e8ba96bd61c895726dd157d67310ae98d1632c9bb9f28e0d7337403c0a0100", + "04000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000000", + ); + const CHIA_CHALLENGE_42_NONTRIVIAL_PROOF_HEX: &str = concat!( + "0000235f6d0bfcbadbd5a0d6619a8611345eb63891876d37150fdef725695ab80c6deef7684c38fe0e086355baf4786fed8a5f843d0b7a62bf1125765b016dfe965b493cfc9bcde723c5299db8db25885d130f9aef4b029f98f42831aaf53e51e3350100", + "0300d2b31e34c399ec49288e3fccb6ebaf0f3fb2e814c7c21e8579c17b5f2600b1a64d9d5b94435084b3458a9343fd1bcd3f0b9e5874556f1ab1529347b54788af1eb9268a5ee888fba85934c81b199a4228a41cb01c10b3195c95b26c17f16ff7020100", + ); + + #[test] + fn prover_matches_chia_known_vector() { + let solution = prove(&[0x42; 32], 100, |_, _| {}).unwrap(); + + assert_eq!( + solution, + crate::domain::decode_hex(CHIA_CHALLENGE_42_PROOF_HEX).unwrap() + ); + assert!(verify(&[0x42; 32], 100, &solution)); + } + + #[test] + fn prover_matches_chia_nontrivial_proof_vector() { + let solution = prove(&[0x42; 32], 300, |_, _| {}).unwrap(); + + assert_eq!( + solution, + crate::domain::decode_hex(CHIA_CHALLENGE_42_NONTRIVIAL_PROOF_HEX).unwrap() + ); + assert!(verify(&[0x42; 32], 300, &solution)); + } + + #[test] + fn proof_is_exactly_two_bqfc_forms() { + let solution = prove(b"iuna-vdf-wire-format", 16, |_, _| {}).unwrap(); + + assert_eq!(solution.len(), SOLUTION_BYTES); + } + + #[test] + fn malformed_and_mismatched_proofs_are_rejected() { + let solution = prove(b"iuna-vdf-negative-test", 300, |_, _| {}).unwrap(); + let mut tampered = solution.clone(); + tampered[50] ^= 1; + + assert!(!verify(b"iuna-vdf-negative-test", 300, &tampered)); + assert!(!verify(b"iuna-vdf-negative-test", 301, &solution)); + assert!(!verify(b"other-seed", 300, &solution)); + assert!(!verify(b"iuna-vdf-negative-test", 300, &solution[..199])); + } + + #[test] + fn non_canonical_special_form_encodings_are_rejected() { + let mut solution = prove(&[0x42; 32], 100, |_, _| {}).unwrap(); + assert_eq!(solution[100], 0x04, "the known proof is the identity"); + + solution[199] = 1; + + assert!(!verify(&[0x42; 32], 100, &solution)); + } +}