commit aa1223b10c7c925699c2b7aca64f886bdd29bec6
parent fb71f0f7fbfcfec715eda75e5a40eda0e9071b0d
Author: Joris Hartog <jorishartog@hotmail.com>
Date: Fri, 21 Aug 2026 13:34:32 +0200
Optimize Rust VDF prover
Diffstat:
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("ient, &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));
+ }
+}