diff --git a/Dockerfile b/Dockerfile index 996c5b4c..9ca05e51 100644 --- a/Dockerfile +++ b/Dockerfile @@ -18,24 +18,13 @@ WORKDIR /app COPY xmss/rust/ xmss/rust/ # Detect the build-stage architecture: legacy builders do not populate TARGETARCH. -# Match make ffi's Haswell/AVX2 baseline on x86_64; leave arm64 flags unchanged. +# Same flags as make ffi: Haswell on x86_64, +aes (PMULL) on aarch64. RUN cd xmss/rust && \ - if [ "$(uname -m)" = "x86_64" ]; then \ - CARGO_ENCODED_RUSTFLAGS="-Ctarget-cpu=haswell" cargo build --profile multisig-release --locked; \ - else \ - cargo build --profile multisig-release --locked; \ - fi - -# Stage leanVM Python sources at the exact checkout path the binary expects. -# The lean_compiler resolves .py files via CARGO_MANIFEST_DIR baked at compile time; -# on arm64 the pre-committed bytecode cache misses and triggers a recompile from source. -# Match the checkout by crate, not by pinned rev: cargo names the rev subdir after -# the leanVM commit, so a hardcoded short hash breaks on every dependency bump. -RUN CHECKOUT_DIR=$(ls -d /root/.cargo/git/checkouts/leanvm-*/*/crates/rec_aggregation | head -1 | sed 's|/crates/rec_aggregation||') && \ - mkdir -p /leanvm-staged && \ - echo "$CHECKOUT_DIR" > /leanvm-staged/.checkout_root && \ - cp -r "$CHECKOUT_DIR/crates/rec_aggregation" /leanvm-staged/rec_aggregation && \ - cp -r "$CHECKOUT_DIR/crates/lean_compiler" /leanvm-staged/lean_compiler + case "$(uname -m)" in \ + x86_64) CARGO_ENCODED_RUSTFLAGS="-Ctarget-cpu=haswell" cargo build --profile multisig-release --locked ;; \ + aarch64) CARGO_ENCODED_RUSTFLAGS="-Ctarget-feature=+aes" cargo build --profile multisig-release --locked ;; \ + *) cargo build --profile multisig-release --locked ;; \ + esac # Copy Go module files for dependency caching COPY go.mod go.sum ./ @@ -68,17 +57,6 @@ LABEL org.opencontainers.image.ref.name=$GIT_BRANCH COPY --from=builder /app/bin/gean /usr/local/bin/ COPY --from=builder /app/bin/keygen /usr/local/bin/ -# leanVM's lean_compiler reads .py files at runtime when the embedded -# cached_bytecode.bin fingerprint doesn't match the build target (arm64 builds -# hit this because the repo's cache is x86-only). Restore the Python sources -# at the exact CARGO_MANIFEST_DIR path baked into the binary at compile time. -COPY --from=builder /leanvm-staged/ /tmp/leanvm-staged/ -RUN CHECKOUT_ROOT=$(cat /tmp/leanvm-staged/.checkout_root) && \ - mkdir -p "$CHECKOUT_ROOT/crates" && \ - cp -r /tmp/leanvm-staged/rec_aggregation "$CHECKOUT_ROOT/crates/" && \ - cp -r /tmp/leanvm-staged/lean_compiler "$CHECKOUT_ROOT/crates/" && \ - rm -rf /tmp/leanvm-staged - # Prove on jemalloc, not glibc malloc. The prover frees its scratch after every # proof, but glibc keeps it: most lands in the main heap, which only shrinks from @@ -94,6 +72,14 @@ RUN apt-get update && apt-get install -y --no-install-recommends libjemalloc2 \ && ldconfig -p | grep -q 'libjemalloc\.so\.2' ENV LD_PRELOAD=libjemalloc.so.2 +# jemalloc purges freed pages only when the arena that freed them is used again, +# unless its background threads are on, and they are off by default. A proof +# spreads gigabytes across the prover's worker threads and then leaves those +# arenas idle, so hundreds of megabytes stayed resident between proofs. Devnet, +# 16 nodes over 10 hours, the only difference being this setting: 268 MB average +# against 729 MB, with identical proof and verification times. +ENV MALLOC_CONF=background_thread:true + # Keep the Go heap tight so the XMSS prover's transient multi-GB proving # peaks (allocated by the Rust arena, invisible to the Go GC) land on free # memory instead of an uncollected heap. Operators can override. diff --git a/Makefile b/Makefile index d7899f7c..242d7e89 100644 --- a/Makefile +++ b/Makefile @@ -16,13 +16,17 @@ LEAN_SPEC_COMMIT_HASH := eca701efeb5931010fe63925cd203c9ee55b2dbc help: ## Show help for each Makefile recipe @grep -E '^[a-zA-Z0-9_-]+:.*?## .*$$' $(MAKEFILE_LIST) | sort | awk 'BEGIN {FS = ":.*?## "}; {printf "\033[36m%-30s\033[0m %s\n", $$1, $$2}' +# leanVM's prover does binary-field arithmetic with carryless multiplication. On x86 the +# Haswell baseline provides PCLMULQDQ; on aarch64 PMULL is gated behind the `aes` target +# feature, which Linux aarch64 targets do not enable by default. Without it the prover +# falls back to scalar code. ffi: ## Build XMSS FFI glue libraries (hashsig-glue + multisig-glue) @cd xmss/rust && \ - if [ "$$(uname -m)" = "x86_64" ]; then \ - CARGO_ENCODED_RUSTFLAGS="-Ctarget-cpu=haswell" cargo build --profile multisig-release --locked; \ - else \ - cargo build --profile multisig-release --locked; \ - fi + case "$$(uname -m)" in \ + x86_64) CARGO_ENCODED_RUSTFLAGS="-Ctarget-cpu=haswell" cargo build --profile multisig-release --locked ;; \ + aarch64|arm64) CARGO_ENCODED_RUSTFLAGS="-Ctarget-feature=+aes" cargo build --profile multisig-release --locked ;; \ + *) cargo build --profile multisig-release --locked ;; \ + esac build: ffi ## Build gean and keygen binaries @mkdir -p bin diff --git a/README.md b/README.md index f229069b..75d4e37f 100644 --- a/README.md +++ b/README.md @@ -75,11 +75,14 @@ via the environment if your budget differs. Gean tracks Lean Consensus devnet-5. Consensus fixtures are generated from `leanSpec@eca701efeb5931010fe63925cd203c9ee55b2dbc`, which pins -`lean-multisig-py` v0.0.9. The XMSS FFI builds against -leanVM `e2592df4e30fdddbbf8ae26a333116c68cec7026`. +`lean-multisig-py` v0.0.9. The XMSS FFI builds against leanVM +`48a904208d682848dac0e18ef8b01ebfc40df9ad`, the BLAKE2s line: keys and proofs +from the earlier Poseidon line do not verify against it, and validator keys must +come from a keygen built on the same rev. `LEAN_SPEC_COMMIT_HASH` in the [`Makefile`](Makefile) is the source of truth for the spec version; the leanVM rev is pinned in +[`xmss/rust/hashsig-glue/Cargo.toml`](xmss/rust/hashsig-glue/Cargo.toml) and [`xmss/rust/multisig-glue/Cargo.toml`](xmss/rust/multisig-glue/Cargo.toml). ## Philosophy diff --git a/internal/blockbuilder/build_test.go b/internal/blockbuilder/build_test.go index a122ffef..97c44948 100644 --- a/internal/blockbuilder/build_test.go +++ b/internal/blockbuilder/build_test.go @@ -508,10 +508,56 @@ func TestPlanAttestationsDoesNotReportSkippedPayloadsWhenFull(t *testing.T) { if err != nil { t.Fatalf("plan attestations: %v", err) } - if len(plan.attestations) != int(types.MaxAttestationsData) { - t.Fatalf("planned attestations=%d, want %d", len(plan.attestations), types.MaxAttestationsData) + if len(plan.attestations) != maxBlockAttestationData { + t.Fatalf("planned attestations=%d, want %d", len(plan.attestations), maxBlockAttestationData) } if len(plan.payloadErrors) != 0 { t.Fatalf("payload errors=%d, want 0", len(plan.payloadErrors)) } } + +// Validators that disagree within a slot cast distinct AttestationData at that slot, and a +// claim group is an (epoch, message) pair, so the block carries them all. The block's own +// slot is no exception: the proposer's signature over the block root is its own group there. +func TestPlanAttestationsKeepsEveryDataAtOneSlot(t *testing.T) { + headState, parentRoot, data, dataRoot := postHeaderVoteInput(t) + root1 := [32]byte{0x11} + + // Same slot as data, different head: a second message at a slot the block already covers. + rival := *data + rival.Head = &types.Checkpoint{Slot: 1, Root: root1} + // The block's own slot, which the proposer's signature also holds. + atBlockSlot := *data + atBlockSlot.Slot = 3 + + plan, err := planAttestations(Input{ + HeadState: headState, + Slot: 3, + ProposerIndex: 0, + ParentRoot: parentRoot, + KnownBlockRoots: RootSet{root1: true, parentRoot: true}, + Payloads: []AttestationPayload{ + {DataRoot: dataRoot, Data: data, Proofs: []*types.SingleMessageAggregate{mockProof([]uint64{0})}}, + {DataRoot: hashAttestationData(t, &rival), Data: &rival, Proofs: []*types.SingleMessageAggregate{mockProof([]uint64{0})}}, + {DataRoot: hashAttestationData(t, &atBlockSlot), Data: &atBlockSlot, Proofs: []*types.SingleMessageAggregate{mockProof([]uint64{0})}}, + }, + }) + if err != nil { + t.Fatalf("plan attestations: %v", err) + } + if len(plan.attestations) != 3 { + t.Fatalf("planned %d attestations, want 3", len(plan.attestations)) + } + if len(plan.payloadErrors) != 0 { + t.Fatalf("payload errors=%d, want 0: %v", len(plan.payloadErrors), plan.payloadErrors) + } + atSlot2 := 0 + for _, att := range plan.attestations { + if att.Data.Slot == 2 { + atSlot2++ + } + } + if atSlot2 != 2 { + t.Fatalf("attestations at slot 2 = %d, want 2", atSlot2) + } +} diff --git a/internal/blockbuilder/plan.go b/internal/blockbuilder/plan.go index 06441c50..bfe8c596 100644 --- a/internal/blockbuilder/plan.go +++ b/internal/blockbuilder/plan.go @@ -64,8 +64,15 @@ func planAttestations(input Input) (planResult, error) { return planner.result(), nil } +// maxBlockAttestationData caps the distinct AttestationData in a block this node +// builds, below the types.MaxAttestationsData that imported blocks are checked +// against. Each one is a child of the block's Type-2 proof, and the proof's time and +// peak memory grow with every child: with 8 it takes 5-10 s, past the proposal +// deadline, and peaks near 9 GB. +const maxBlockAttestationData = 3 + func (p *planner) run() error { - for len(p.attestations) < int(types.MaxAttestationsData) { + for len(p.attestations) < maxBlockAttestationData { if !p.runRound() { return nil } @@ -73,7 +80,7 @@ func (p *planner) run() error { if err != nil { return err } - if len(p.attestations) >= int(types.MaxAttestationsData) { + if len(p.attestations) >= maxBlockAttestationData { p.full = true return nil } @@ -88,7 +95,7 @@ func (p *planner) run() error { func (p *planner) runRound() bool { added := false for _, payload := range p.payloads { - if len(p.attestations) >= int(types.MaxAttestationsData) { + if len(p.attestations) >= maxBlockAttestationData { return added } if p.tryPayload(payload) { diff --git a/internal/blockprocessor/block_bench_test.go b/internal/blockprocessor/block_bench_test.go index 3d47f458..ffb5bf67 100644 --- a/internal/blockprocessor/block_bench_test.go +++ b/internal/blockprocessor/block_bench_test.go @@ -111,7 +111,7 @@ func buildBenchSignedBlock(n int) (*types.SignedBlock, error) { } bits := types.BitlistFromIndices([]uint64{0}) atts[i] = &types.AggregatedAttestation{AggregationBits: bits, Data: data} - inputs = append(inputs, xmss.Type1Input{Pubkeys: []xmss.CPubKey{benchPubkey}, Proof: proof}) + inputs = append(inputs, xmss.Type1Input{Pubkeys: []xmss.CPubKey{benchPubkey}, Proof: proof, Message: root, Slot: uint32(data.Slot)}) } block := &types.Block{ @@ -136,18 +136,9 @@ func buildBenchSignedBlock(n int) (*types.SignedBlock, error) { if err != nil { return nil, err } - proposerProof, err := xmss.AggregateSignatures( - []xmss.CPubKey{benchPubkey}, - []xmss.CSig{signature}, - blockRoot, - blockSlot, - ) - xmss.FreeSignature(signature) - if err != nil { - return nil, err - } - inputs = append(inputs, xmss.Type1Input{Pubkeys: []xmss.CPubKey{benchPubkey}, Proof: proposerProof}) - proof, err := xmss.MergeType1Proofs(inputs) + defer xmss.FreeSignature(signature) + proposer := xmss.RawSignature{Pubkey: benchPubkey, Signature: signature, Message: blockRoot, Slot: blockSlot} + proof, err := xmss.MergeType1Proofs(inputs, []xmss.RawSignature{proposer}) if err != nil { return nil, err } diff --git a/internal/metrics/histograms.go b/internal/metrics/histograms.go index 1136dbd6..9bef4e51 100644 --- a/internal/metrics/histograms.go +++ b/internal/metrics/histograms.go @@ -121,6 +121,15 @@ var ( Help: "Elapsed time between clock ticks in seconds", Buckets: []float64{0.4, 0.6, 0.75, 0.8, 0.805, 0.81, 0.815, 0.82, 0.825, 0.85, 0.9, 1.0, 1.2, 1.6}, }) + // metricTickPhase is where in its interval each tick's duties start. An + // aligned clock reads near zero; a constant offset means the tick schedule + // is shifted and every duty runs that much late. Dispatch-loop delay before + // the tick is handled counts too, since that is when the duties really run. + metricTickPhase = promauto.NewHistogram(prometheus.HistogramOpts{ + Name: "lean_tick_phase_seconds", + Help: "Offset of each tick into its interval, measured when the tick is handled", + Buckets: []float64{0.001, 0.0025, 0.005, 0.01, 0.025, 0.05, 0.1, 0.2, 0.4, 0.6, 0.8}, + }) // metricDispatchEventDuration times each case of the dispatch select, so a // slow handler can be attributed rather than only observed as a late tick. // Buckets run well past a slot: the point is to size a stall, and the diff --git a/internal/metrics/observe.go b/internal/metrics/observe.go index 20dcc9f0..c6bf627a 100644 --- a/internal/metrics/observe.go +++ b/internal/metrics/observe.go @@ -25,6 +25,7 @@ func ObserveForkChoiceReorgDepth(depth float64) { func ObserveTickIntervalDuration(seconds float64) { observeNonNegative(metricTickIntervalDuration, seconds) } +func ObserveTickPhase(seconds float64) { observeNonNegative(metricTickPhase, seconds) } // ObserveDispatchEvent records how long one dispatch-loop event took. func ObserveDispatchEvent(event string, seconds float64) { diff --git a/internal/node/clock.go b/internal/node/clock.go index 1ccbe3cf..52e905e4 100644 --- a/internal/node/clock.go +++ b/internal/node/clock.go @@ -16,6 +16,21 @@ func (e *Engine) currentInterval(timestampMs uint64) uint64 { return types.CurrentInterval(e.Store.Config().GenesisTime, timestampMs) } +// claimInterval reports whether the interval containing timestampMs has not +// had its duties run yet, and marks it as run. When the dispatch loop falls +// behind, two ticks can be handled inside one interval; running its duties a +// second time would sign a second attestation for the slot. Before genesis +// millisIntoSlot is 0, so each pre-genesis tick claims its own timestamp and +// none can shadow the genesis interval. +func (e *Engine) claimInterval(timestampMs uint64) bool { + start := timestampMs - e.millisIntoSlot(timestampMs)%types.MillisecondsPerInterval + if start <= e.lastIntervalStartMs { + return false + } + e.lastIntervalStartMs = start + return true +} + func (e *Engine) millisIntoSlot(timestampMs uint64) uint64 { if e == nil || e.Store == nil { return 0 diff --git a/internal/node/engine.go b/internal/node/engine.go index 806e0137..aba55171 100644 --- a/internal/node/engine.go +++ b/internal/node/engine.go @@ -91,6 +91,10 @@ type Engine struct { // goroutine precisely so it still reports while the dispatch loop is blocked. lastTickMs atomic.Int64 + // lastIntervalStartMs is the start of the interval whose duties onTick + // last ran; see claimInterval. + lastIntervalStartMs uint64 + warnedMissingJustified [32]byte // maxSeenGossipSlot is the highest plausible slot heard on gossip, whether @@ -176,12 +180,9 @@ func (e *Engine) WaitForStorageWorkers() { func (e *Engine) Run(ctx context.Context) { e.initMetrics() - ticker := time.NewTicker(types.MillisecondsPerInterval * time.Millisecond) - defer ticker.Stop() - e.startWorkers(ctx) logger.Info(logger.Node, "started") e.onTick() - e.dispatch(ctx, ticker.C) + e.dispatch(ctx, alignedTicks(ctx, e.Store.Config().GenesisTime)) } diff --git a/internal/node/proposal.go b/internal/node/proposal.go index ea3fe60c..9cb0bf33 100644 --- a/internal/node/proposal.go +++ b/internal/node/proposal.go @@ -208,16 +208,15 @@ func (e *Engine) mergeBlockProof( proposerKey *xmss.ValidatorKeyPair, proposerSignature [types.SignatureSize]byte, ) ([]byte, error) { - return e.mergeBlockProofWithProvers(block, attestationProofs, proposerKey, proposerSignature, xmss.AggregateSignatures, xmss.MergeType1Proofs) + return e.mergeBlockProofWithProver(block, attestationProofs, proposerKey, proposerSignature, xmss.MergeType1Proofs) } -func (e *Engine) mergeBlockProofWithProvers( +func (e *Engine) mergeBlockProofWithProver( block *types.Block, attestationProofs []*types.SingleMessageAggregate, proposerKey *xmss.ValidatorKeyPair, proposerSignature [types.SignatureSize]byte, - wrap func([]xmss.CPubKey, []xmss.CSig, [32]byte, uint32) ([]byte, error), - merge func([]xmss.Type1Input) ([]byte, error), + merge func([]xmss.Type1Input, []xmss.RawSignature) ([]byte, error), ) ([]byte, error) { if block == nil || block.Body == nil || len(block.Body.Attestations) != len(attestationProofs) { return nil, fmt.Errorf("attestation proof count mismatch") @@ -230,7 +229,7 @@ func (e *Engine) mergeBlockProofWithProvers( return nil, fmt.Errorf("parent state missing") } - inputs := make([]xmss.Type1Input, 0, len(attestationProofs)+1) + inputs := make([]xmss.Type1Input, 0, len(attestationProofs)) for i, proof := range attestationProofs { if proof == nil { return nil, fmt.Errorf("attestation proof %d missing", i) @@ -246,7 +245,12 @@ func (e *Engine) mergeBlockProofWithProvers( } keys = append(keys, key) } - inputs = append(inputs, xmss.Type1Input{Pubkeys: keys, Proof: proof.Proof}) + data := block.Body.Attestations[i].Data + root, err := data.HashTreeRoot() + if err != nil { + return nil, fmt.Errorf("attestation %d data root: %w", i, err) + } + inputs = append(inputs, xmss.Type1Input{Pubkeys: keys, Proof: proof.Proof, Message: root, Slot: uint32(data.Slot)}) } signature, err := xmss.ParseSignature(proposerSignature[:]) @@ -261,26 +265,17 @@ func (e *Engine) mergeBlockProofWithProvers( if e.Store.Head() != block.ParentRoot { return nil, errStaleProposal } - wrapStart := time.Now() - proposerProof, err := wrap( - []xmss.CPubKey{proposerKey.PublicKey()}, - []xmss.CSig{signature}, - blockRoot, - uint32(block.Slot), - ) - metrics.ObserveProposalStageDuration("signature_proof", time.Since(wrapStart).Seconds()) - if err != nil { - return nil, err - } - inputs = append(inputs, xmss.Type1Input{ - Pubkeys: []xmss.CPubKey{proposerKey.PublicKey()}, - Proof: proposerProof, - }) - if e.Store.Head() != block.ParentRoot { - return nil, errStaleProposal + // The proposer's signature goes into the merge raw. Proving it alone first and + // merging that proof costs a whole extra proof, the most expensive part of a + // proposal after the merge itself. + proposer := xmss.RawSignature{ + Pubkey: proposerKey.PublicKey(), + Signature: signature, + Message: blockRoot, + Slot: uint32(block.Slot), } mergeStart := time.Now() - proof, err := merge(inputs) + proof, err := merge(inputs, []xmss.RawSignature{proposer}) metrics.ObserveProposalStageDuration("merge", time.Since(mergeStart).Seconds()) return proof, err } diff --git a/internal/node/proposal_test.go b/internal/node/proposal_test.go index 418ae969..dd8b7266 100644 --- a/internal/node/proposal_test.go +++ b/internal/node/proposal_test.go @@ -303,17 +303,15 @@ func TestProposalSigningErrorRetainsDuty(t *testing.T) { } } -func TestBlockProofStopsBetweenStagesWhenParentChanges(t *testing.T) { +func TestBlockProofStopsWhenParentChanges(t *testing.T) { for _, tc := range []struct { - name string - staleBefore, staleAfterWrap, wrapFails, mergeFails bool - wantWrap, wantMerge int + name string + staleBefore, mergeFails bool + wantMerge int }{ {name: "already_stale", staleBefore: true}, - {name: "stale_during_signature_proof", staleAfterWrap: true, wantWrap: 1}, - {name: "unchanged_parent", wantWrap: 1, wantMerge: 1}, - {name: "signature_proof_error", wrapFails: true, wantWrap: 1}, - {name: "merge_error", mergeFails: true, wantWrap: 1, wantMerge: 1}, + {name: "unchanged_parent", wantMerge: 1}, + {name: "merge_error", mergeFails: true, wantMerge: 1}, } { t.Run(tc.name, func(t *testing.T) { e := proposalTestEngine(t) @@ -325,41 +323,30 @@ func TestBlockProofStopsBetweenStagesWhenParentChanges(t *testing.T) { if tc.staleBefore { e.Store.SetHead([32]byte{99}) } - wrapCalls, mergeCalls := 0, 0 + mergeCalls := 0 proofErr := errors.New("test prover failure") - wrap := func(pks []xmss.CPubKey, sigs []xmss.CSig, message [32]byte, slot uint32) ([]byte, error) { - wrapCalls++ - if len(pks) != 1 || len(sigs) != 1 || message != root || slot != 1 { - t.Fatal("proposer binding changed") - } - if tc.staleAfterWrap { - e.Store.SetHead([32]byte{99}) - } - if tc.wrapFails { - return nil, proofErr - } - return []byte{7}, nil - } - merge := func(inputs []xmss.Type1Input) ([]byte, error) { + merge := func(inputs []xmss.Type1Input, raw []xmss.RawSignature) ([]byte, error) { mergeCalls++ - if len(inputs) != 1 || len(inputs[0].Proof) != 1 || inputs[0].Proof[0] != 7 { - t.Fatal("proposer proof missing from final merge") + // The proposer's signature enters the merge raw, bound to the block root + // at the block's slot, rather than as a proof of its own. + if len(inputs) != 0 || len(raw) != 1 || raw[0].Message != root || raw[0].Slot != 1 { + t.Fatal("proposer signature not merged raw with its block binding") } if tc.mergeFails { return nil, proofErr } return []byte{8}, nil } - proof, err := e.mergeBlockProofWithProvers(block, nil, nil, [types.SignatureSize]byte{}, wrap, merge) - if wrapCalls != tc.wantWrap || mergeCalls != tc.wantMerge { - t.Fatalf("wrap=%d merge=%d", wrapCalls, mergeCalls) + proof, err := e.mergeBlockProofWithProver(block, nil, nil, [types.SignatureSize]byte{}, merge) + if mergeCalls != tc.wantMerge { + t.Fatalf("merge=%d", mergeCalls) } switch { - case tc.staleBefore || tc.staleAfterWrap: + case tc.staleBefore: if !errors.Is(err, errStaleProposal) || proof != nil { t.Fatalf("expected stale failure: %v", err) } - case tc.wrapFails || tc.mergeFails: + case tc.mergeFails: if !errors.Is(err, proofErr) { t.Fatalf("lost proof error: %v", err) } diff --git a/internal/node/recovery.go b/internal/node/recovery.go index 25e5ada5..e30d42a6 100644 --- a/internal/node/recovery.go +++ b/internal/node/recovery.go @@ -105,7 +105,7 @@ func (e *Engine) recoverBlockProofs(ctx context.Context, signedBlock *types.Sign if headState == nil || headState.LatestJustified == nil { return } - pubkeys, err := e.blockProofPubkeys(block, state) + pubkeys, bindings, err := e.blockProofClaims(block, state) if err != nil { return } @@ -143,7 +143,7 @@ func (e *Engine) recoverBlockProofs(ctx context.Context, signedBlock *types.Sign return } started := time.Now() - proof, err := xmss.SplitType2Proof(signedBlock.Proof.Proof, pubkeys, candidate.root) + proof, err := xmss.SplitType2Proof(signedBlock.Proof.Proof, pubkeys, bindings, candidate.root) var recovered *types.SingleMessageAggregate if err == nil { recovered = &types.SingleMessageAggregate{ @@ -234,30 +234,45 @@ func coversParticipants(proof *types.SingleMessageAggregate, participants []byte return true } -func (e *Engine) blockProofPubkeys(block *types.Block, state *types.State) ([][]xmss.CPubKey, error) { +// blockProofClaims returns the signer groups of a block's Type-2 proof and what each +// signed, in the order block verification binds them: one group per attestation, then +// the proposer over the block root. +func (e *Engine) blockProofClaims(block *types.Block, state *types.State) ([][]xmss.CPubKey, []xmss.MessageBinding, error) { groups := make([][]xmss.CPubKey, 0, len(block.Body.Attestations)+1) + bindings := make([]xmss.MessageBinding, 0, len(block.Body.Attestations)+1) for _, att := range block.Body.Attestations { keys := make([]xmss.CPubKey, 0, types.BitlistCount(att.AggregationBits)) for _, index := range types.BitlistIndices(att.AggregationBits) { if index >= uint64(len(state.Validators)) || state.Validators[index] == nil { - return nil, fmt.Errorf("validator %d out of range", index) + return nil, nil, fmt.Errorf("validator %d out of range", index) } key, err := e.Store.PubKeyCache.Get(state.Validators[index].AttestationPubkey) if err != nil { - return nil, err + return nil, nil, err } keys = append(keys, key) } + root, err := att.Data.HashTreeRoot() + if err != nil { + return nil, nil, err + } groups = append(groups, keys) + bindings = append(bindings, xmss.MessageBinding{Message: root, Slot: uint32(att.Data.Slot)}) } if block.ProposerIndex >= uint64(len(state.Validators)) || state.Validators[block.ProposerIndex] == nil { - return nil, fmt.Errorf("proposer %d out of range", block.ProposerIndex) + return nil, nil, fmt.Errorf("proposer %d out of range", block.ProposerIndex) } key, err := e.Store.PubKeyCache.Get(state.Validators[block.ProposerIndex].ProposalPubkey) if err != nil { - return nil, err + return nil, nil, err + } + blockRoot, err := block.HashTreeRoot() + if err != nil { + return nil, nil, err } - return append(groups, []xmss.CPubKey{key}), nil + groups = append(groups, []xmss.CPubKey{key}) + bindings = append(bindings, xmss.MessageBinding{Message: blockRoot, Slot: uint32(block.Slot)}) + return groups, bindings, nil } func localCoverage(entries ...*store.PayloadEntry) map[uint64]bool { diff --git a/internal/node/tick.go b/internal/node/tick.go index fe3aa390..4e23c662 100644 --- a/internal/node/tick.go +++ b/internal/node/tick.go @@ -4,6 +4,7 @@ import ( "time" "github.com/geanlabs/gean/internal/aggregation" + "github.com/geanlabs/gean/internal/logger" "github.com/geanlabs/gean/internal/metrics" "github.com/geanlabs/gean/internal/store" "github.com/geanlabs/gean/internal/types" @@ -29,6 +30,18 @@ func (e *Engine) onTick() { currentSlot := e.currentSlot(timestampMs) currentInterval := e.currentInterval(timestampMs) + // The startup tick lands wherever the process began, so only scheduled + // ticks say anything about the clock's phase. + if !firstTick { + phaseMs := e.millisIntoSlot(timestampMs) % types.MillisecondsPerInterval + metrics.ObserveTickPhase(float64(phaseMs) / 1000) + } + + if !e.claimInterval(timestampMs) { + logger.Warn(logger.Node, "tick skipped: interval already handled slot=%d interval=%d", currentSlot, currentInterval) + return + } + metrics.SetCurrentSlot(currentSlot) e.updateSyncStatus(currentSlot) diff --git a/internal/node/ticks.go b/internal/node/ticks.go new file mode 100644 index 00000000..ab59a9b9 --- /dev/null +++ b/internal/node/ticks.go @@ -0,0 +1,66 @@ +package node + +import ( + "context" + "time" + + "github.com/geanlabs/gean/internal/types" +) + +// alignedTicks delivers one tick per interval boundary, measured from genesis. +// A time.Ticker would keep the period but take its phase from process start. +// +// Each boundary is delivered at most once and never before the wall clock +// reaches it. Delivery is a non-blocking send into a one-slot buffer, as with +// time.Ticker: a tick missed while dispatch is busy is dropped, not queued. +// A tick handled late can still land in an interval already run; onTick's +// claimInterval drops it. +func alignedTicks(ctx context.Context, genesisTime uint64) <-chan time.Time { + ch := make(chan time.Time, 1) + go func() { + timer := time.NewTimer(time.Hour) + defer timer.Stop() + var lastMs uint64 + for { + targetMs := nextTickAt(genesisTime, uint64(time.Now().UnixMilli()), lastMs) + // Timers measure a duration on the monotonic clock; the boundary is + // on the wall clock. Re-check after every wake so a slewed or stepped + // wall clock can never release a tick early. + for { + wait := time.Until(time.UnixMilli(int64(targetMs))) + if wait <= 0 { + break + } + timer.Reset(wait) + select { + case <-ctx.Done(): + return + case <-timer.C: + } + } + lastMs = targetMs + select { + case ch <- time.Now(): + default: + } + } + }() + return ch +} + +// nextTickAt returns the wall-clock millisecond of the next tick to deliver: +// the first interval boundary after nowMs, and never one at or before lastMs, +// the boundary delivered previously (zero before the first). Holding to lastMs +// keeps delivery monotonic when the wall clock steps backwards. If genesis is +// unrepresentable it falls back to a plain interval from now, the old +// free-running behaviour, rather than stopping the clock. +func nextTickAt(genesisTime, nowMs, lastMs uint64) uint64 { + target, ok := types.NextIntervalBoundaryMs(genesisTime, nowMs) + if !ok { + return nowMs + types.MillisecondsPerInterval + } + if lastMs != 0 && target <= lastMs { + return lastMs + types.MillisecondsPerInterval + } + return target +} diff --git a/internal/node/ticks_test.go b/internal/node/ticks_test.go new file mode 100644 index 00000000..3e37a392 --- /dev/null +++ b/internal/node/ticks_test.go @@ -0,0 +1,119 @@ +package node + +import ( + "context" + "testing" + "time" + + "github.com/geanlabs/gean/internal/types" +) + +func TestNextTickAt(t *testing.T) { + const gt = 1700000000 + const gtMs = gt * 1000 + const iv = uint64(types.MillisecondsPerInterval) + tests := []struct { + name string + nowMs uint64 + lastMs uint64 + want uint64 + }{ + {"first_tick_before_genesis_waits_for_genesis", gtMs - 3000, 0, gtMs}, + {"first_tick_mid_interval_goes_to_next_boundary", gtMs + 10*iv + 790, 0, gtMs + 11*iv}, + {"woken_exactly_on_boundary_moves_a_full_interval", gtMs + 11*iv, gtMs + 11*iv, gtMs + 12*iv}, + {"normal_progression", gtMs + 11*iv + 3, gtMs + 11*iv, gtMs + 12*iv}, + // A wall clock stepped back behind the last delivered boundary must not + // deliver that boundary, or any earlier one, a second time. + {"clock_stepped_back_never_repeats_a_boundary", gtMs + 9*iv + 100, gtMs + 11*iv, gtMs + 12*iv}, + {"clock_stepped_forward_skips_ahead", gtMs + 20*iv + 5, gtMs + 11*iv, gtMs + 21*iv}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := nextTickAt(gt, tt.nowMs, tt.lastMs); got != tt.want { + t.Fatalf("nextTickAt(%d, %d, %d) = %d, want %d", gt, tt.nowMs, tt.lastMs, got, tt.want) + } + }) + } +} + +func TestNextTickAtFallsBackWhenGenesisOverflows(t *testing.T) { + const nowMs = 1_700_000_000_000 + if got := nextTickAt(^uint64(0)/1000+1, nowMs, 0); got != nowMs+types.MillisecondsPerInterval { + t.Fatalf("overflowing genesis: got %d, want a plain interval from now", got) + } +} + +// TestAlignedTicksLandOnIntervalBoundary checks the property the free-running +// ticker lacked: the tick arrives at the start of an interval measured from +// genesis, whatever moment the source was started at. +func TestAlignedTicksLandOnIntervalBoundary(t *testing.T) { + genesis := uint64(time.Now().Unix()) - 10 + // Start the source partway into an interval, the position a node lands in + // after an arbitrary restart. + nowMs := uint64(time.Now().UnixMilli()) + next, _ := types.NextIntervalBoundaryMs(genesis, nowMs) + if wait := time.Until(time.UnixMilli(int64(next))); wait > 0 { + time.Sleep(wait + 300*time.Millisecond) + } + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + ticks := alignedTicks(ctx, genesis) + + select { + case tick := <-ticks: + phaseMs := types.MillisIntoSlot(genesis, uint64(tick.UnixMilli())) % types.MillisecondsPerInterval + // A tick released early would read just below a full interval. + if phaseMs > 50 { + t.Fatalf("tick landed %d ms into its interval; want it at the boundary", phaseMs) + } + case <-time.After(2 * types.MillisecondsPerInterval * time.Millisecond): + t.Fatal("no tick within two intervals") + } +} + +// TestClaimIntervalRunsEachIntervalOnce covers a tick handled late, after the +// next boundary's tick was already queued: both land in one interval, and only +// the first may run its duties. +func TestClaimIntervalRunsEachIntervalOnce(t *testing.T) { + e := makeTestEngine() + gMs := e.Store.Config().GenesisTime * 1000 + iv := uint64(types.MillisecondsPerInterval) + steps := []struct { + name string + atMs uint64 + claim bool + }{ + {"first_tick_in_interval", gMs + 5*iv + 10, true}, + {"second_tick_same_interval", gMs + 5*iv + 700, false}, + {"next_interval", gMs + 6*iv + 3, true}, + {"clock_stepped_back_into_a_handled_interval", gMs + 5*iv + 790, false}, + {"skipped_interval_then_later_one", gMs + 8*iv + 1, true}, + } + for _, s := range steps { + if got := e.claimInterval(s.atMs); got != s.claim { + t.Fatalf("%s: claimInterval(%d) = %v, want %v", s.name, s.atMs, got, s.claim) + } + } +} + +// TestClaimIntervalBeforeGenesis: a node starts before genesis, so pre-genesis +// ticks must neither collide with each other nor shadow the genesis interval. +func TestClaimIntervalBeforeGenesis(t *testing.T) { + e := makeTestEngine() + gMs := e.Store.Config().GenesisTime * 1000 + for _, s := range []struct { + name string + atMs uint64 + claim bool + }{ + {"startup_before_genesis", gMs - 5000, true}, + {"later_pre_genesis_tick", gMs - 4200, true}, + {"genesis_tick", gMs, true}, + {"same_interval_as_genesis", gMs + 100, false}, + } { + if got := e.claimInterval(s.atMs); got != s.claim { + t.Fatalf("%s: claimInterval(%d) = %v, want %v", s.name, s.atMs, got, s.claim) + } + } +} diff --git a/internal/spectests/poseidon_test.go b/internal/spectests/poseidon_test.go deleted file mode 100644 index 2f22e998..00000000 --- a/internal/spectests/poseidon_test.go +++ /dev/null @@ -1,57 +0,0 @@ -//go:build spectests - -package spectests - -import ( - "encoding/json" - "strconv" - "testing" - - "github.com/geanlabs/gean/xmss" -) - -type poseidonCase struct { - Width int `json:"width"` - InputState []string `json:"inputState"` - OutputState []string `json:"outputState"` -} - -func TestSpecPoseidonPermutation(t *testing.T) { - walkFixtures(t, "../../leanSpec/fixtures/consensus/poseidon_permutation", func(t *testing.T, raw []byte) { - var fixture map[string]poseidonCase - if err := json.Unmarshal(raw, &fixture); err != nil { - t.Fatalf("unmarshal: %v", err) - } - for name, tc := range fixture { - tc := tc - t.Run(name, func(t *testing.T) { - state := parseFieldElements(t, tc.InputState) - if err := xmss.Poseidon2Permute(state); err != nil { - t.Fatalf("permute width %d: %v", tc.Width, err) - } - want := parseFieldElements(t, tc.OutputState) - if len(state) != len(want) { - t.Fatalf("length mismatch: got %d want %d", len(state), len(want)) - } - for i := range want { - if state[i] != want[i] { - t.Fatalf("element %d: got %d want %d", i, state[i], want[i]) - } - } - }) - } - }) -} - -func parseFieldElements(t *testing.T, vals []string) []uint32 { - t.Helper() - out := make([]uint32, len(vals)) - for i, v := range vals { - n, err := strconv.ParseUint(v, 10, 32) - if err != nil { - t.Fatalf("parse field element %q: %v", v, err) - } - out[i] = uint32(n) - } - return out -} diff --git a/internal/types/clock.go b/internal/types/clock.go index 1cdad241..dca63b13 100644 --- a/internal/types/clock.go +++ b/internal/types/clock.go @@ -27,6 +27,29 @@ func CurrentInterval(genesisTime, currentTimeMs uint64) uint64 { return MillisIntoSlot(genesisTime, currentTimeMs) / MillisecondsPerInterval } +// NextIntervalBoundaryMs returns the unix-millisecond time of the first +// interval boundary strictly after currentTimeMs. Exactly on a boundary, that +// is the following one, a full interval away. Before genesis it is genesis +// itself. ok is false when the boundary does not fit in a uint64. +func NextIntervalBoundaryMs(genesisTime, currentTimeMs uint64) (uint64, bool) { + genesisMs, ok := unixMillis(genesisTime) + if !ok { + return 0, false + } + if currentTimeMs < genesisMs { + return genesisMs, true + } + elapsedIntervals := (currentTimeMs - genesisMs) / MillisecondsPerInterval + if elapsedIntervals >= ^uint64(0)/MillisecondsPerInterval { + return 0, false + } + offset := (elapsedIntervals + 1) * MillisecondsPerInterval + if offset > ^uint64(0)-genesisMs { + return 0, false + } + return genesisMs + offset, true +} + func TotalIntervals(genesisTime, currentTimeMs uint64) uint64 { genesisMs, ok := unixMillis(genesisTime) if !ok { diff --git a/internal/types/clock_test.go b/internal/types/clock_test.go index 34641217..c9b9f468 100644 --- a/internal/types/clock_test.go +++ b/internal/types/clock_test.go @@ -115,3 +115,40 @@ func TestClockDerivationsHandleOverflow(t *testing.T) { t.Fatalf("IntervalsFromSlot overflow=%d, want max", got) } } + +func TestNextIntervalBoundaryMs(t *testing.T) { + const gt = 1700000000 + const gtMs = gt * 1000 + const iv = uint64(MillisecondsPerInterval) + tests := []struct { + name string + currentMs uint64 + want uint64 + }{ + {"long_before_genesis", gtMs - 5000, gtMs}, + {"1ms_before_genesis", gtMs - 1, gtMs}, + {"at_genesis_is_a_full_interval_away", gtMs, gtMs + iv}, + {"1ms_after_genesis", gtMs + 1, gtMs + iv}, + {"1ms_before_boundary", gtMs + iv - 1, gtMs + iv}, + {"on_boundary_is_a_full_interval_away", gtMs + iv, gtMs + 2*iv}, + {"into_last_interval_rolls_to_next_slot", gtMs + 4*iv + 10, gtMs + uint64(MillisecondsPerSlot)}, + {"deep_into_the_chain", gtMs + 57456*uint64(MillisecondsPerSlot) + 790, gtMs + 57456*uint64(MillisecondsPerSlot) + iv}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, ok := NextIntervalBoundaryMs(gt, tt.currentMs) + if !ok || got != tt.want { + t.Fatalf("NextIntervalBoundaryMs(%d, %d) = (%d, %v), want (%d, true)", gt, tt.currentMs, got, ok, tt.want) + } + }) + } +} + +func TestNextIntervalBoundaryMsOverflow(t *testing.T) { + if _, ok := NextIntervalBoundaryMs(^uint64(0)/1000+1, 0); ok { + t.Fatal("genesis beyond uint64 milliseconds must report !ok") + } + if _, ok := NextIntervalBoundaryMs(0, ^uint64(0)); ok { + t.Fatal("a boundary past uint64 must report !ok") + } +} diff --git a/xmss/block_roundtrip_test.go b/xmss/block_roundtrip_test.go index b240ca08..e213f661 100644 --- a/xmss/block_roundtrip_test.go +++ b/xmss/block_roundtrip_test.go @@ -45,11 +45,8 @@ func TestProposerSigThroughBlockSSZ(t *testing.T) { } defer FreePublicKey(cpk) - proof, err := AggregateSignatures([]CPubKey{cpk}, []CSig{csig}, blockRoot, 1) - if err != nil { - t.Fatalf("aggregate FAILED: %v", err) - } - merged, err := MergeType1Proofs([]Type1Input{{Pubkeys: []CPubKey{cpk}, Proof: proof}}) + // A block without attestations: its proof is the proposer's raw signature alone. + merged, err := MergeType1Proofs(nil, []RawSignature{{Pubkey: cpk, Signature: csig, Message: blockRoot, Slot: 1}}) if err != nil { t.Fatalf("merge FAILED: %v", err) } diff --git a/xmss/ffi.go b/xmss/ffi.go index 65e1adf4..e72a69d7 100644 --- a/xmss/ffi.go +++ b/xmss/ffi.go @@ -68,12 +68,17 @@ package xmss // int32_t xmss_merge_type_1_to_type_2( // const uint8_t* const* proof_ptrs, const size_t* proof_lens, // const PublicKey* const* pubkeys, const size_t* pubkey_counts, -// size_t count, size_t log_inv_rate, +// const uint8_t* message_hashes, const uint32_t* message_slots, +// size_t count, +// const PublicKey* const* raw_pub_keys, const Signature* const* raw_signatures, +// const uint8_t* raw_message_hashes, const uint32_t* raw_slots, size_t num_raw, +// size_t log_inv_rate, // uint8_t* out_buf, size_t out_cap, size_t* out_written); // // int32_t xmss_split_type_2_by_message( // const uint8_t* proof, size_t proof_len, // const PublicKey* const* pubkeys, const size_t* pubkey_counts, +// const uint8_t* message_hashes, const uint32_t* message_slots, // size_t count, const uint8_t* target_message, size_t log_inv_rate, // uint8_t* out_buf, size_t out_cap, size_t* out_written); // @@ -81,9 +86,6 @@ package xmss // const uint8_t* proof, size_t proof_len, // const PublicKey* const* pubkeys, const size_t* pubkey_counts, // size_t count, const uint8_t* message_hashes, const uint32_t* message_slots); -// -// int32_t poseidon_permute_kb16(uint32_t* state, size_t len); -// int32_t poseidon_permute_kb24(uint32_t* state, size_t len); import "C" import ( @@ -121,30 +123,8 @@ var ( ErrMalformedChildProof = errors.New("malformed child proof") ErrMalformedRawInput = errors.New("malformed raw signature input") ErrSetupFailed = errors.New("XMSS setup failed") - ErrPoseidonWidth = errors.New("poseidon permutation width must be 16 or 24") - ErrPoseidonPermute = errors.New("poseidon permutation failed") ) -// Poseidon2Permute applies the KoalaBear Poseidon permutation in place. The -// state holds canonical field-element values; width must be 16 or 24. -func Poseidon2Permute(state []uint32) error { - if len(state) != 16 && len(state) != 24 { - return ErrPoseidonWidth - } - ptr := (*C.uint32_t)(unsafe.Pointer(&state[0])) - n := C.size_t(len(state)) - var status C.int32_t - if len(state) == 16 { - status = C.poseidon_permute_kb16(ptr, n) - } else { - status = C.poseidon_permute_kb24(ptr, n) - } - if status != 0 { - return ErrPoseidonPermute - } - return nil -} - var ( proverOnce sync.Once verifierOnce sync.Once @@ -170,11 +150,13 @@ func SetProverArena(enabled bool) { func EnsureProverReady() error { proverOnce.Do(func() { var status C.int32_t - if proverArena.Load() { - status = C.xmss_setup_prover() - } else { - status = C.xmss_setup_prover_without_arena() - } + onProverThread(func() { + if proverArena.Load() { + status = C.xmss_setup_prover() + } else { + status = C.xmss_setup_prover_without_arena() + } + }) if status != 0 { proverErr = ErrSetupFailed } @@ -312,27 +294,30 @@ func AggregateWithChildren( defer putProofBuf(bufPtr) buf := *bufPtr var written C.size_t - status := C.xmss_aggregate_type_1( - rawPkPtr, - rawSigPtr, - C.size_t(numRaw), - childAllPkPtr, - childNumKeysPtr, - childProofPtrsPtr, - childProofLensPtr, - C.size_t(numChildren), - (*C.uint8_t)(unsafe.Pointer(&message[0])), - C.uint32_t(slot), - C.size_t(LogInvRate), - (*C.uint8_t)(unsafe.Pointer(&buf[0])), - C.size_t(len(buf)), - &written, - ) + var status C.int32_t + onProverThread(func() { + status = C.xmss_aggregate_type_1( + rawPkPtr, + rawSigPtr, + C.size_t(numRaw), + childAllPkPtr, + childNumKeysPtr, + childProofPtrsPtr, + childProofLensPtr, + C.size_t(numChildren), + (*C.uint8_t)(unsafe.Pointer(&message[0])), + C.uint32_t(slot), + C.size_t(LogInvRate), + (*C.uint8_t)(unsafe.Pointer(&buf[0])), + C.size_t(len(buf)), + &written, + ) + }) if status != 0 { if status == -2 || int(written) > MaxProofSize { return nil, ErrProofTooBig } - return nil, ErrAggregationFailed + return nil, proofFailure(status) } if written == 0 { return nil, ErrSerializationFailed @@ -428,9 +413,13 @@ func VerifyAggregatedSignature( return nil } +// Type1Input is one Type-1 proof to merge, with the message and slot it proves. Proof +// bytes carry no claims of their own, so the merge has to be told what each one proves. type Type1Input struct { Pubkeys []CPubKey Proof []byte + Message [MessageLength]byte + Slot uint32 } type MessageBinding struct { @@ -438,8 +427,19 @@ type MessageBinding struct { Slot uint32 } -func MergeType1Proofs(inputs []Type1Input) ([]byte, error) { - if len(inputs) == 0 { +// RawSignature is a signature merged into a Type-2 as it is, without first being proved +// on its own. +type RawSignature struct { + Pubkey CPubKey + Signature CSig + Message [MessageLength]byte + Slot uint32 +} + +// MergeType1Proofs merges Type-1 proofs, and raw signatures folded in directly, into one +// Type-2. Either list may be empty, not both. +func MergeType1Proofs(inputs []Type1Input, raw []RawSignature) ([]byte, error) { + if len(inputs) == 0 && len(raw) == 0 { return nil, ErrEmptyInput } if err := EnsureProverReady(); err != nil { @@ -449,57 +449,112 @@ func MergeType1Proofs(inputs []Type1Input) ([]byte, error) { var pinner runtime.Pinner defer pinner.Unpin() - proofPtrs := make([]*C.uint8_t, len(inputs)) - proofLens := make([]C.size_t, len(inputs)) - keyCounts := make([]C.size_t, len(inputs)) - var keys []*C.PublicKey - for i, input := range inputs { - if len(input.Proof) == 0 || len(input.Proof) > MaxProofSize { - return nil, ErrMalformedChildProof + var ( + proofPtrsPtr **C.uint8_t + proofLensPtr *C.size_t + keysPtr **C.PublicKey + countsPtr *C.size_t + hashesPtr *C.uint8_t + slotsPtr *C.uint32_t + ) + if len(inputs) > 0 { + proofPtrs := make([]*C.uint8_t, len(inputs)) + proofLens := make([]C.size_t, len(inputs)) + keyCounts := make([]C.size_t, len(inputs)) + bindings := make([]MessageBinding, len(inputs)) + var keys []*C.PublicKey + for i, input := range inputs { + bindings[i] = MessageBinding{Message: input.Message, Slot: input.Slot} + if len(input.Proof) == 0 || len(input.Proof) > MaxProofSize { + return nil, ErrMalformedChildProof + } + if err := validatePublicKeys(input.Pubkeys); err != nil { + return nil, err + } + pinner.Pin(&input.Proof[0]) + proofPtrs[i] = (*C.uint8_t)(unsafe.Pointer(&input.Proof[0])) + proofLens[i] = C.size_t(len(input.Proof)) + keyCounts[i] = C.size_t(len(input.Pubkeys)) + for _, key := range input.Pubkeys { + keys = append(keys, (*C.PublicKey)(key)) + } } - if err := validatePublicKeys(input.Pubkeys); err != nil { - return nil, err + if len(keys) == 0 { + return nil, ErrEmptyInput } - pinner.Pin(&input.Proof[0]) - proofPtrs[i] = (*C.uint8_t)(unsafe.Pointer(&input.Proof[0])) - proofLens[i] = C.size_t(len(input.Proof)) - keyCounts[i] = C.size_t(len(input.Pubkeys)) - for _, key := range input.Pubkeys { - keys = append(keys, (*C.PublicKey)(key)) + hashes, slots := flattenBindings(bindings) + pinner.Pin(&proofPtrs[0]) + pinner.Pin(&keys[0]) + proofPtrsPtr = (**C.uint8_t)(unsafe.Pointer(&proofPtrs[0])) + proofLensPtr = (*C.size_t)(unsafe.Pointer(&proofLens[0])) + keysPtr = (**C.PublicKey)(unsafe.Pointer(&keys[0])) + countsPtr = (*C.size_t)(unsafe.Pointer(&keyCounts[0])) + hashesPtr = (*C.uint8_t)(unsafe.Pointer(&hashes[0])) + slotsPtr = (*C.uint32_t)(unsafe.Pointer(&slots[0])) + } + + var ( + rawKeysPtr **C.PublicKey + rawSigsPtr **C.Signature + rawHashesPtr *C.uint8_t + rawSlotsPtr *C.uint32_t + ) + if len(raw) > 0 { + rawKeys := make([]*C.PublicKey, len(raw)) + rawSigs := make([]*C.Signature, len(raw)) + bindings := make([]MessageBinding, len(raw)) + for i, r := range raw { + if r.Pubkey == nil || r.Signature == nil { + return nil, fmt.Errorf("%w: raw signature %d", ErrMalformedRawInput, i) + } + rawKeys[i] = (*C.PublicKey)(r.Pubkey) + rawSigs[i] = (*C.Signature)(r.Signature) + bindings[i] = MessageBinding{Message: r.Message, Slot: r.Slot} } + hashes, slots := flattenBindings(bindings) + pinner.Pin(&rawKeys[0]) + pinner.Pin(&rawSigs[0]) + rawKeysPtr = (**C.PublicKey)(unsafe.Pointer(&rawKeys[0])) + rawSigsPtr = (**C.Signature)(unsafe.Pointer(&rawSigs[0])) + rawHashesPtr = (*C.uint8_t)(unsafe.Pointer(&hashes[0])) + rawSlotsPtr = (*C.uint32_t)(unsafe.Pointer(&slots[0])) } - if len(keys) == 0 { - return nil, ErrEmptyInput - } - pinner.Pin(&proofPtrs[0]) - pinner.Pin(&keys[0]) bufPtr := getProofBuf() defer putProofBuf(bufPtr) buf := *bufPtr var written C.size_t - status := C.xmss_merge_type_1_to_type_2( - (**C.uint8_t)(unsafe.Pointer(&proofPtrs[0])), - (*C.size_t)(unsafe.Pointer(&proofLens[0])), - (**C.PublicKey)(unsafe.Pointer(&keys[0])), - (*C.size_t)(unsafe.Pointer(&keyCounts[0])), - C.size_t(len(inputs)), - C.size_t(LogInvRate), - (*C.uint8_t)(unsafe.Pointer(&buf[0])), - C.size_t(len(buf)), - &written, - ) + var status C.int32_t + onProverThread(func() { + status = C.xmss_merge_type_1_to_type_2( + proofPtrsPtr, proofLensPtr, keysPtr, countsPtr, hashesPtr, slotsPtr, + C.size_t(len(inputs)), + rawKeysPtr, rawSigsPtr, rawHashesPtr, rawSlotsPtr, + C.size_t(len(raw)), + C.size_t(LogInvRate), + (*C.uint8_t)(unsafe.Pointer(&buf[0])), + C.size_t(len(buf)), + &written, + ) + }) return proofResult(status, written, buf) } +// SplitType2Proof re-proves the group of a Type-2 proof that signed target as a +// standalone Type-1. pubkeys and bindings describe every group of the Type-2, in the same +// order, since the proof bytes carry no claims of their own. func SplitType2Proof( proof []byte, pubkeys [][]CPubKey, + bindings []MessageBinding, target [MessageLength]byte, ) ([]byte, error) { if len(proof) == 0 || len(pubkeys) == 0 { return nil, ErrEmptyInput } + if len(pubkeys) != len(bindings) { + return nil, ErrCountMismatch + } if len(proof) > MaxProofSize { return nil, ErrProofTooBig } @@ -511,6 +566,7 @@ func SplitType2Proof( if err != nil { return nil, err } + hashes, slots := flattenBindings(bindings) var pinner runtime.Pinner defer pinner.Unpin() pinner.Pin(&proof[0]) @@ -520,18 +576,23 @@ func SplitType2Proof( defer putProofBuf(bufPtr) buf := *bufPtr var written C.size_t - status := C.xmss_split_type_2_by_message( - (*C.uint8_t)(unsafe.Pointer(&proof[0])), - C.size_t(len(proof)), - (**C.PublicKey)(unsafe.Pointer(&keys[0])), - (*C.size_t)(unsafe.Pointer(&counts[0])), - C.size_t(len(pubkeys)), - (*C.uint8_t)(unsafe.Pointer(&target[0])), - C.size_t(LogInvRate), - (*C.uint8_t)(unsafe.Pointer(&buf[0])), - C.size_t(len(buf)), - &written, - ) + var status C.int32_t + onProverThread(func() { + status = C.xmss_split_type_2_by_message( + (*C.uint8_t)(unsafe.Pointer(&proof[0])), + C.size_t(len(proof)), + (**C.PublicKey)(unsafe.Pointer(&keys[0])), + (*C.size_t)(unsafe.Pointer(&counts[0])), + (*C.uint8_t)(unsafe.Pointer(&hashes[0])), + (*C.uint32_t)(unsafe.Pointer(&slots[0])), + C.size_t(len(pubkeys)), + (*C.uint8_t)(unsafe.Pointer(&target[0])), + C.size_t(LogInvRate), + (*C.uint8_t)(unsafe.Pointer(&buf[0])), + C.size_t(len(buf)), + &written, + ) + }) return proofResult(status, written, buf) } @@ -554,12 +615,7 @@ func VerifyType2Proof( return err } - hashes := make([]byte, 0, len(bindings)*MessageLength) - slots := make([]C.uint32_t, len(bindings)) - for i, binding := range bindings { - hashes = append(hashes, binding.Message[:]...) - slots[i] = C.uint32_t(binding.Slot) - } + hashes, slots := flattenBindings(bindings) var pinner runtime.Pinner defer pinner.Unpin() @@ -599,12 +655,48 @@ func flattenPublicKeys(groups [][]CPubKey) ([]*C.PublicKey, []C.size_t, error) { return keys, counts, nil } +// flattenBindings lays bindings out as the C side reads them: consecutive 32-byte +// messages, and the slots in a parallel array. +func flattenBindings(bindings []MessageBinding) ([]byte, []C.uint32_t) { + hashes := make([]byte, 0, len(bindings)*MessageLength) + slots := make([]C.uint32_t, len(bindings)) + for i, binding := range bindings { + hashes = append(hashes, binding.Message[:]...) + slots[i] = C.uint32_t(binding.Slot) + } + return hashes, slots +} + +// proofFailures names the status codes the proving calls return (Failure in +// multisig-glue), so a proof that could not be built says why. +var proofFailures = map[C.int32_t]string{ + -3: "an input proof does not decode against its claims", + -5: "split target is not a group of the proof", + -6: "prover panicked", + -10: "an input proof does not verify", + -11: "malformed raw signature", + -12: "too many claim groups", + -13: "too many inputs or signers", + -14: "an input proof is too large", + -15: "declared claim is not covered by the inputs", + -16: "nothing to prove", + -17: "invalid proof rate", + -18: "invalid data-availability input", +} + +func proofFailure(status C.int32_t) error { + if reason, ok := proofFailures[status]; ok { + return fmt.Errorf("%w: %s", ErrAggregationFailed, reason) + } + return ErrAggregationFailed +} + func proofResult(status C.int32_t, written C.size_t, buf []byte) ([]byte, error) { if status != 0 { if status == -2 || int(written) > MaxProofSize { return nil, ErrProofTooBig } - return nil, ErrAggregationFailed + return nil, proofFailure(status) } if written == 0 { return nil, ErrSerializationFailed diff --git a/xmss/prover_thread.go b/xmss/prover_thread.go new file mode 100644 index 00000000..4b4df12c --- /dev/null +++ b/xmss/prover_thread.go @@ -0,0 +1,35 @@ +package xmss + +import ( + "runtime" + "sync" +) + +// Every prover call runs on one OS thread. With the arena on, leanVM gives each +// thread that proves its own region and keeps that thread's peak resident for the +// life of the process. A cgo call runs on whichever thread the scheduler picked +// for the calling goroutine, so proofs issued from different goroutines would each +// leave a peak behind on a different thread until the process runs out of memory. +// Proofs already run one at a time, so a single thread costs no parallelism. +var ( + proverJobs = make(chan func()) + proverThreadOnce sync.Once +) + +// onProverThread runs f on the prover thread and returns once it has finished. +func onProverThread(f func()) { + proverThreadOnce.Do(func() { go runProverThread() }) + done := make(chan struct{}) + proverJobs <- func() { + defer close(done) + f() + } + <-done +} + +func runProverThread() { + runtime.LockOSThread() + for job := range proverJobs { + job() + } +} diff --git a/xmss/prover_thread_test.go b/xmss/prover_thread_test.go new file mode 100644 index 00000000..5d6a91d2 --- /dev/null +++ b/xmss/prover_thread_test.go @@ -0,0 +1,58 @@ +//go:build linux + +package xmss + +import ( + "runtime" + "sync" + "syscall" + "testing" + "time" +) + +// TestProverJobsShareOneThread submits jobs from callers locked to distinct OS +// threads, the way gean's workers reach the prover from different threads, and +// requires every job to run on one thread that none of the callers own. +func TestProverJobsShareOneThread(t *testing.T) { + const callers = 8 + var mu sync.Mutex + callerTids := map[int]bool{} + jobTids := map[int]bool{} + + var wg sync.WaitGroup + for range callers { + wg.Add(1) + go func() { + defer wg.Done() + runtime.LockOSThread() + defer runtime.UnlockOSThread() + mu.Lock() + callerTids[syscall.Gettid()] = true + mu.Unlock() + for range 5 { + onProverThread(func() { + first := syscall.Gettid() + // Give the scheduler a chance to move an unlocked goroutine. + time.Sleep(time.Millisecond) + mu.Lock() + jobTids[first] = true + jobTids[syscall.Gettid()] = true + mu.Unlock() + }) + } + }() + } + wg.Wait() + + if len(callerTids) != callers { + t.Fatalf("callers ran on %d threads, want %d distinct", len(callerTids), callers) + } + if len(jobTids) != 1 { + t.Fatalf("prover jobs ran on %d OS threads, want 1", len(jobTids)) + } + for tid := range jobTids { + if callerTids[tid] { + t.Fatalf("prover job ran on caller thread %d", tid) + } + } +} diff --git a/xmss/rust/Cargo.lock b/xmss/rust/Cargo.lock index e7878ec6..6764b140 100644 --- a/xmss/rust/Cargo.lock +++ b/xmss/rust/Cargo.lock @@ -11,16 +11,6 @@ dependencies = [ "memchr", ] -[[package]] -name = "air" -version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" -dependencies = [ - "field", - "koala-bear", - "poly", -] - [[package]] name = "alloy-primitives" version = "1.6.1" @@ -423,24 +413,6 @@ version = "1.5.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f2032f911046de80f0a198e0901378627c33f59ea0ac00e363d481118bd70a53" -[[package]] -name = "backend" -version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" -dependencies = [ - "air", - "fiat-shamir", - "field", - "koala-bear", - "parallel", - "poly", - "sumcheck", - "symetric", - "utils", - "whir", - "zk-alloc", -] - [[package]] name = "base16ct" version = "0.2.0" @@ -459,6 +431,15 @@ version = "1.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2af50177e190e07a26ab74f8b1efbfe2ef87da2116221318cb1c2e82baf7de06" +[[package]] +name = "bincode" +version = "1.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b1f45e9417d87227c7a56d22e471c6206462cba514c7590c09aff4cf6d1ddcad" +dependencies = [ + "serde", +] + [[package]] name = "bitcoin-consensus-encoding" version = "1.1.0" @@ -606,17 +587,6 @@ dependencies = [ "multisig-glue", ] -[[package]] -name = "chacha20" -version = "0.10.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81" -dependencies = [ - "cfg-if", - "cpufeatures 0.3.0", - "rand_core 0.10.1", -] - [[package]] name = "chrono" version = "0.4.45" @@ -1070,27 +1040,13 @@ dependencies = [ ] [[package]] -name = "fiat-shamir" +name = "fiat_shamir" version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" +source = "git+https://github.com/leanEthereum/leanVM.git?rev=48a904208d682848dac0e18ef8b01ebfc40df9ad#48a904208d682848dac0e18ef8b01ebfc40df9ad" dependencies = [ - "field", - "koala-bear", "parallel", + "primitives", "serde", - "symetric", - "utils", -] - -[[package]] -name = "field" -version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" -dependencies = [ - "paste", - "rand 0.10.2", - "serde", - "utils", ] [[package]] @@ -1121,6 +1077,18 @@ dependencies = [ "static_assertions", ] +[[package]] +name = "flock" +version = "0.1.0" +source = "git+https://github.com/leanEthereum/leanVM.git?rev=48a904208d682848dac0e18ef8b01ebfc40df9ad#48a904208d682848dac0e18ef8b01ebfc40df9ad" +dependencies = [ + "fiat_shamir", + "parallel", + "pcs", + "primitives", + "zk_alloc", +] + [[package]] name = "foldhash" version = "0.2.0" @@ -1187,22 +1155,10 @@ checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" dependencies = [ "cfg-if", "libc", - "r-efi 5.3.0", + "r-efi", "wasip2", ] -[[package]] -name = "getrandom" -version = "0.4.3" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "300e883d756b2e4ec94e02791f39b04b522276138852cfc41d9fb7e904106099" -dependencies = [ - "cfg-if", - "libc", - "r-efi 6.0.0", - "rand_core 0.10.1", -] - [[package]] name = "group" version = "0.13.0" @@ -1244,10 +1200,9 @@ dependencies = [ name = "hashsig-glue" version = "0.1.0" dependencies = [ - "ethereum_ssz", + "leanvm", "postcard", "sha2 0.9.9", - "xmss", ] [[package]] @@ -1356,25 +1311,6 @@ dependencies = [ "syn 2.0.119", ] -[[package]] -name = "include_dir" -version = "0.7.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "923d117408f1e49d914f1a379a309cffe4f18c05cf4e3d12e613a15fc81bd0dd" -dependencies = [ - "include_dir_macros", -] - -[[package]] -name = "include_dir_macros" -version = "0.7.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7cab85a7ed0bd5f0e76d93846e0147172bed2e2d3f859bcc33a8d9699cad1a75" -dependencies = [ - "proc-macro2", - "quote", -] - [[package]] name = "indexmap" version = "1.9.3" @@ -1534,18 +1470,6 @@ dependencies = [ "sha3-asm", ] -[[package]] -name = "koala-bear" -version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" -dependencies = [ - "field", - "paste", - "rand 0.10.2", - "serde", - "utils", -] - [[package]] name = "konst" version = "0.2.20" @@ -1568,58 +1492,55 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" [[package]] -name = "lean-multisig" +name = "lean_compiler" version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" +source = "git+https://github.com/leanEthereum/leanVM.git?rev=48a904208d682848dac0e18ef8b01ebfc40df9ad#48a904208d682848dac0e18ef8b01ebfc40df9ad" dependencies = [ - "backend", - "clap", "lean_vm", - "rec_aggregation", - "serde_json", - "sub_protocols", - "system-info", - "xmss", - "zk-alloc", + "primitives", ] [[package]] -name = "lean_compiler" +name = "lean_da" version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" +source = "git+https://github.com/leanEthereum/leanVM.git?rev=48a904208d682848dac0e18ef8b01ebfc40df9ad#48a904208d682848dac0e18ef8b01ebfc40df9ad" dependencies = [ - "backend", - "include_dir", - "lean_vm", - "pest", - "pest_derive", - "sub_protocols", - "xmss", + "fiat_shamir", + "parallel", + "pcs", + "primitives", + "serde", + "tracing", ] [[package]] -name = "lean_prover" +name = "lean_vm" version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" +source = "git+https://github.com/leanEthereum/leanVM.git?rev=48a904208d682848dac0e18ef8b01ebfc40df9ad#48a904208d682848dac0e18ef8b01ebfc40df9ad" dependencies = [ - "backend", - "lean_compiler", - "lean_vm", - "rand 0.10.2", - "serde", - "sub_protocols", + "fiat_shamir", + "flock", + "parallel", + "pcs", + "primitives", "tracing", - "xmss", + "zk_alloc", ] [[package]] -name = "lean_vm" +name = "leanvm" version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" +source = "git+https://github.com/leanEthereum/leanVM.git?rev=48a904208d682848dac0e18ef8b01ebfc40df9ad#48a904208d682848dac0e18ef8b01ebfc40df9ad" dependencies = [ - "backend", - "tracing", + "clap", + "lean_da", + "lean_vm", + "primitives", + "rand 0.9.5", + "rec_aggregation", + "sphincs", "xmss", + "zk_alloc", ] [[package]] @@ -1668,9 +1589,7 @@ checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" name = "multisig-glue" version = "0.1.0" dependencies = [ - "backend", - "lean-multisig", - "rec_aggregation", + "leanvm", ] [[package]] @@ -1717,31 +1636,6 @@ dependencies = [ "libm", ] -[[package]] -name = "objc2" -version = "0.6.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3a12a8ed07aefc768292f076dc3ac8c48f3781c8f2d5851dd3d98950e8c5a89f" -dependencies = [ - "objc2-encode", -] - -[[package]] -name = "objc2-encode" -version = "4.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ef25abbcd74fb2609453eb695bd2f860d389e457f67dc17cafc8b8cbc89d0c33" - -[[package]] -name = "objc2-foundation" -version = "0.3.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e3e0adef53c21f888deb4fa59fc59f7eb17404926ee8a6f59f5df0fd7f9f3272" -dependencies = [ - "bitflags 2.13.1", - "objc2", -] - [[package]] name = "once_cell" version = "1.21.4" @@ -1763,9 +1657,9 @@ checksum = "c08d65885ee38876c4f86fa503fb49d7b507c2b62552df7c70b2fce627e06381" [[package]] name = "parallel" version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" +source = "git+https://github.com/leanEthereum/leanVM.git?rev=48a904208d682848dac0e18ef8b01ebfc40df9ad#48a904208d682848dac0e18ef8b01ebfc40df9ad" dependencies = [ - "system-info", + "libc", ] [[package]] @@ -1803,45 +1697,26 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "57c0d7b74b563b49d38dae00a0c37d4d6de9b432382b2892f0574ddcae73fd0a" [[package]] -name = "pest" -version = "2.8.8" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7df728be843c7070fab6ab7c328c4e9e9d78e23bf749c0669c86ee7ebfa050a2" -dependencies = [ - "memchr", - "ucd-trie", -] - -[[package]] -name = "pest_derive" -version = "2.8.8" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9e2dd6fc3b26b3462ee188aac870f5a41d398f1cd5e2408d16531bd71c9591fd" -dependencies = [ - "pest", - "pest_generator", -] - -[[package]] -name = "pest_generator" -version = "2.8.8" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6a7a9205cfb6f596a9e8b689c0a15f9ceb7a1aafae7aaf788150ac65b29975b6" +name = "pcs" +version = "0.1.0" +source = "git+https://github.com/leanEthereum/leanVM.git?rev=48a904208d682848dac0e18ef8b01ebfc40df9ad#48a904208d682848dac0e18ef8b01ebfc40df9ad" dependencies = [ - "pest", - "pest_meta", - "proc-macro2", - "quote", - "syn 2.0.119", + "fiat_shamir", + "parallel", + "primitives", + "serde", + "tracing", + "zk_alloc", ] [[package]] -name = "pest_meta" +name = "pest" version = "2.8.8" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "85abd351c0de1e8384fc791a0737111a350394937e92b956b743dac12429f57c" +checksum = "7df728be843c7070fab6ab7c328c4e9e9d78e23bf749c0669c86ee7ebfa050a2" dependencies = [ - "pest", + "memchr", + "ucd-trie", ] [[package]] @@ -1860,21 +1735,6 @@ dependencies = [ "spki", ] -[[package]] -name = "poly" -version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" -dependencies = [ - "field", - "koala-bear", - "parallel", - "rand 0.10.2", - "serde", - "system-info", - "utils", - "zk-alloc", -] - [[package]] name = "portable-atomic" version = "1.15.0" @@ -1929,6 +1789,19 @@ dependencies = [ "uint", ] +[[package]] +name = "primitives" +version = "0.1.0" +source = "git+https://github.com/leanEthereum/leanVM.git?rev=48a904208d682848dac0e18ef8b01ebfc40df9ad#48a904208d682848dac0e18ef8b01ebfc40df9ad" +dependencies = [ + "libc", + "parallel", + "serde", + "tracing-forest", + "tracing-subscriber", + "zk_alloc", +] + [[package]] name = "proc-macro-crate" version = "3.5.0" @@ -1977,12 +1850,6 @@ version = "5.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" -[[package]] -name = "r-efi" -version = "6.0.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" - [[package]] name = "radium" version = "0.7.0" @@ -2011,17 +1878,6 @@ dependencies = [ "serde", ] -[[package]] -name = "rand" -version = "0.10.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80" -dependencies = [ - "chacha20", - "getrandom 0.4.3", - "rand_core 0.10.1", -] - [[package]] name = "rand_chacha" version = "0.3.1" @@ -2061,12 +1917,6 @@ dependencies = [ "serde", ] -[[package]] -name = "rand_core" -version = "0.10.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" - [[package]] name = "rand_xorshift" version = "0.4.0" @@ -2088,19 +1938,19 @@ dependencies = [ [[package]] name = "rec_aggregation" version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" +source = "git+https://github.com/leanEthereum/leanVM.git?rev=48a904208d682848dac0e18ef8b01ebfc40df9ad#48a904208d682848dac0e18ef8b01ebfc40df9ad" dependencies = [ - "backend", - "include_dir", + "bincode", + "flock", "lean_compiler", - "lean_prover", + "lean_da", "lean_vm", - "objc2", - "objc2-foundation", "parallel", - "postcard", + "pcs", + "primitives", + "rand 0.9.5", "serde", - "sub_protocols", + "sphincs", "tracing", "xmss", ] @@ -2465,6 +2315,17 @@ version = "1.15.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8ed6a63f02c8539c91a8685a86f4099661ba3da017932f6ebbea6de3f0fa7c90" +[[package]] +name = "sphincs" +version = "0.1.0" +source = "git+https://github.com/leanEthereum/leanVM.git?rev=48a904208d682848dac0e18ef8b01ebfc40df9ad#48a904208d682848dac0e18ef8b01ebfc40df9ad" +dependencies = [ + "parallel", + "primitives", + "rand 0.9.5", + "serde", +] + [[package]] name = "spin" version = "0.9.9" @@ -2502,48 +2363,12 @@ version = "0.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" -[[package]] -name = "sub_protocols" -version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" -dependencies = [ - "backend", - "lean_vm", - "tracing", -] - [[package]] name = "subtle" version = "2.6.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" -[[package]] -name = "sumcheck" -version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" -dependencies = [ - "air", - "fiat-shamir", - "field", - "koala-bear", - "parallel", - "poly", - "tracing", - "zk-alloc", -] - -[[package]] -name = "symetric" -version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" -dependencies = [ - "field", - "koala-bear", - "parallel", - "zk-alloc", -] - [[package]] name = "syn" version = "1.0.109" @@ -2577,14 +2402,6 @@ dependencies = [ "unicode-ident", ] -[[package]] -name = "system-info" -version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" -dependencies = [ - "libc", -] - [[package]] name = "tap" version = "1.0.1" @@ -2823,17 +2640,6 @@ version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" -[[package]] -name = "utils" -version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" -dependencies = [ - "parallel", - "serde", - "tracing-forest", - "tracing-subscriber", -] - [[package]] name = "valuable" version = "0.1.1" @@ -2906,25 +2712,6 @@ dependencies = [ "unicode-ident", ] -[[package]] -name = "whir" -version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" -dependencies = [ - "fiat-shamir", - "field", - "koala-bear", - "parallel", - "poly", - "rand 0.10.2", - "sumcheck", - "symetric", - "system-info", - "tracing", - "utils", - "zk-alloc", -] - [[package]] name = "winapi" version = "0.3.9" @@ -3042,12 +2829,12 @@ dependencies = [ [[package]] name = "xmss" version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" +source = "git+https://github.com/leanEthereum/leanVM.git?rev=48a904208d682848dac0e18ef8b01ebfc40df9ad#48a904208d682848dac0e18ef8b01ebfc40df9ad" dependencies = [ - "backend", "ethereum_ssz", - "postcard", - "rand 0.10.2", + "parallel", + "primitives", + "rand 0.9.5", "serde", ] @@ -3092,13 +2879,11 @@ dependencies = [ ] [[package]] -name = "zk-alloc" +name = "zk_alloc" version = "0.1.0" -source = "git+https://github.com/leanEthereum/leanVM.git?rev=a5909d18647de6aed38640c098d9177fab2bf36a#a5909d18647de6aed38640c098d9177fab2bf36a" +source = "git+https://github.com/leanEthereum/leanVM.git?rev=48a904208d682848dac0e18ef8b01ebfc40df9ad#48a904208d682848dac0e18ef8b01ebfc40df9ad" dependencies = [ "libc", - "parallel", - "system-info", ] [[package]] diff --git a/xmss/rust/hashsig-glue/Cargo.toml b/xmss/rust/hashsig-glue/Cargo.toml index 50359d31..7f276205 100644 --- a/xmss/rust/hashsig-glue/Cargo.toml +++ b/xmss/rust/hashsig-glue/Cargo.toml @@ -5,16 +5,10 @@ edition = "2021" [dependencies] sha2 = "0.9" -# XMSS is now internalized in leanVM (the standalone leanSig dependency is gone). -# The single-signature scheme lives in leanVM's `xmss` crate; aggregation lives in -# `multisig-glue`. Both pin the SAME leanVM rev — keep them in lockstep. -# -# Pinned to a5909d18 to match ethlambda's interop devnet rev; must stay in -# lockstep with multisig-glue's lean-multisig pin (same leanVM rev). -xmss = { git = "https://github.com/leanEthereum/leanVM.git", rev = "a5909d18647de6aed38640c098d9177fab2bf36a" } -# ethereum_ssz 0.10 — must match leanVM's ssz version so the Encode/Decode impls -# on XmssPublicKey / XmssSignature are the same trait. -ssz = { package = "ethereum_ssz", version = "0.10" } +# XMSS comes from leanVM's root `leanvm` crate, which also re-exports the SSZ traits +# its key and signature types implement. multisig-glue depends on the same crate: +# keep both on one rev, since keys, signatures and aggregates only verify within it. +leanvm = { git = "https://github.com/leanEthereum/leanVM.git", rev = "48a904208d682848dac0e18ef8b01ebfc40df9ad" } # XmssSecretKey is serde-only (SSZ deliberately excluded upstream); persist it # with postcard. The secret key never goes on the wire, so the format is local. postcard = { version = "1.1.3", features = ["alloc"] } diff --git a/xmss/rust/hashsig-glue/src/lib.rs b/xmss/rust/hashsig-glue/src/lib.rs index e04178c8..80ab83a2 100644 --- a/xmss/rust/hashsig-glue/src/lib.rs +++ b/xmss/rust/hashsig-glue/src/lib.rs @@ -1,17 +1,18 @@ +use leanvm::xmss::{ + key_gen_from_seed, sign, verify, Decode, Encode, Epoch, XmssPublicKey, XmssSecretKey, + XmssSignature, MESSAGE_LEN, PUB_KEY_SSZ_LEN, SIGNATURE_SSZ_LEN, +}; use sha2::{Digest, Sha256}; -use ssz::{Decode, Encode}; use std::ffi::CStr; use std::os::raw::c_char; use std::ptr; use std::slice; -use xmss::{ - xmss_key_gen_from_seed, xmss_sign, xmss_verify, XmssPublicKey, XmssSecretKey, XmssSignature, - MESSAGE_LEN_BYTES, -}; -// leanVM-internalized XMSS: a single fixed instantiation (V=42, base 8, log-lifetime 32, -// KoalaBear/Poseidon2). SSZ public key = 32 bytes, signature = 1208 bytes. Keys/sigs from -// leanSig's Dim46 aborting scheme are NOT interoperable with this and must be regenerated. +// leanVM XMSS over BLAKE2s: one fixed instantiation (V=42, base 8, log-lifetime 32). Keys and +// signatures from the earlier Poseidon line do not verify here and must be regenerated. + +// The Go SSZ types hard-code these sizes (internal/types/constants.go, xmss.MessageLength). +const _: () = assert!(PUB_KEY_SSZ_LEN == 32 && SIGNATURE_SSZ_LEN == 1208 && MESSAGE_LEN == 32); #[repr(C)] pub struct PrivateKey { @@ -51,8 +52,11 @@ pub unsafe extern "C" fn hashsig_keypair_generate( hasher.update(seed_phrase.as_bytes()); let seed: [u8; 32] = hasher.finalize().into(); - match xmss_key_gen_from_seed(seed, activation_epoch as u64, num_active_epochs as u64) { - Ok((public_key, private_key)) => Box::into_raw(Box::new(KeyPair { + let Some((epoch_start, epoch_end)) = epoch_range(activation_epoch, num_active_epochs) else { + return ptr::null_mut(); + }; + match key_gen_from_seed(seed, epoch_start, epoch_end) { + Ok((private_key, public_key)) => Box::into_raw(Box::new(KeyPair { public_key: PublicKey { inner: public_key }, private_key: PrivateKey { inner: private_key }, })), @@ -60,6 +64,13 @@ pub unsafe extern "C" fn hashsig_keypair_generate( } } +/// leanVM takes an inclusive epoch range; this ABI passes a first epoch and a count. +fn epoch_range(activation_epoch: usize, num_active_epochs: usize) -> Option<(Epoch, Epoch)> { + let start = Epoch::try_from(activation_epoch).ok()?; + let span = Epoch::try_from(num_active_epochs.checked_sub(1)?).ok()?; + Some((start, start.checked_add(span)?)) +} + /// Reconstruct a key pair from its persisted parts: the secret key is postcard (serde), the /// public key is SSZ. The two encodings differ because upstream persists the secret key with /// serde and deliberately excludes it from SSZ. @@ -163,12 +174,12 @@ pub unsafe extern "C" fn hashsig_sign( } unsafe { let private_key_ref = &*private_key; - let message_slice = slice::from_raw_parts(message_ptr, MESSAGE_LEN_BYTES); - let message_array: &[u8; MESSAGE_LEN_BYTES] = match message_slice.try_into() { + let message_slice = slice::from_raw_parts(message_ptr, MESSAGE_LEN); + let message_array: &[u8; MESSAGE_LEN] = match message_slice.try_into() { Ok(arr) => arr, Err(_) => return ptr::null_mut(), }; - match xmss_sign(&private_key_ref.inner, epoch, message_array) { + match sign(&private_key_ref.inner, message_array, epoch) { Ok(sig) => Box::into_raw(Box::new(Signature { inner: sig })), Err(_) => ptr::null_mut(), } @@ -215,16 +226,16 @@ pub unsafe extern "C" fn hashsig_verify( unsafe { let public_key_ref = &*public_key; let signature_ref = &*signature; - let message_slice = slice::from_raw_parts(message_ptr, MESSAGE_LEN_BYTES); - let message_array: &[u8; MESSAGE_LEN_BYTES] = match message_slice.try_into() { + let message_slice = slice::from_raw_parts(message_ptr, MESSAGE_LEN); + let message_array: &[u8; MESSAGE_LEN] = match message_slice.try_into() { Ok(arr) => arr, Err(_) => return -1, }; - match xmss_verify( + match verify( &public_key_ref.inner, - epoch, message_array, &signature_ref.inner, + epoch, ) { Ok(()) => 1, Err(_) => 0, @@ -234,7 +245,7 @@ pub unsafe extern "C" fn hashsig_verify( #[no_mangle] pub extern "C" fn hashsig_message_length() -> usize { - MESSAGE_LEN_BYTES + MESSAGE_LEN } #[no_mangle] @@ -320,8 +331,8 @@ pub unsafe extern "C" fn hashsig_verify_ssz( unsafe { let pk_data = slice::from_raw_parts(pubkey_bytes, pubkey_len); let sig_data = slice::from_raw_parts(signature_bytes, signature_len); - let msg_data = slice::from_raw_parts(message, MESSAGE_LEN_BYTES); - let message_array: &[u8; MESSAGE_LEN_BYTES] = match msg_data.try_into() { + let msg_data = slice::from_raw_parts(message, MESSAGE_LEN); + let message_array: &[u8; MESSAGE_LEN] = match msg_data.try_into() { Ok(arr) => arr, Err(_) => return -1, }; @@ -333,7 +344,7 @@ pub unsafe extern "C" fn hashsig_verify_ssz( Ok(sig) => sig, Err(_) => return -1, }; - match xmss_verify(&pk, epoch, message_array, &sig) { + match verify(&pk, message_array, &sig, epoch) { Ok(()) => 1, Err(_) => 0, } diff --git a/xmss/rust/multisig-glue/Cargo.toml b/xmss/rust/multisig-glue/Cargo.toml index a59dd440..881890af 100644 --- a/xmss/rust/multisig-glue/Cargo.toml +++ b/xmss/rust/multisig-glue/Cargo.toml @@ -4,21 +4,9 @@ version = "0.1.0" edition = "2021" [dependencies] -# leanVM main HEAD. The root crate re-exports the aggregation API, Xmss -# key/signature types, and setup_prover/setup_verifier (which own the internal -# arena allocator; no global_allocator override). XMSS is now internalized here -# (leanSig is gone), so this rev must match hashsig-glue's `xmss` pin exactly. -# -# Pinned to a5909d18 to match ethlambda's interop devnet rev (same aggregate wire -# format) and to pick up setup_prover_without_arena for the low-memory prover path. -# Aggregate wire format changed vs e2592df (compress/decompress_without_pubkeys -# → to_bytes/from_bytes_without_pubkeys; info.without_pubkeys → info.core). -lean-multisig = { git = "https://github.com/leanEthereum/leanVM.git", rev = "a5909d18647de6aed38640c098d9177fab2bf36a" } -# split-by-message is not re-exported by the root crate; pull it from the -# aggregation crate directly (same git rev, identical re-exported types). -rec_aggregation = { git = "https://github.com/leanEthereum/leanVM.git", rev = "a5909d18647de6aed38640c098d9177fab2bf36a" } -# Poseidon permutation primitives for the spec test-vector FFI (poseidon_permute_*). -backend = { git = "https://github.com/leanEthereum/leanVM.git", rev = "a5909d18647de6aed38640c098d9177fab2bf36a" } +# Same crate and rev as hashsig-glue: an aggregate verifies only against keys and +# signatures from the leanVM rev that produced it. +leanvm = { git = "https://github.com/leanEthereum/leanVM.git", rev = "48a904208d682848dac0e18ef8b01ebfc40df9ad" } [lib] # Built as an rlib so its `#[no_mangle] pub extern "C"` symbols can be diff --git a/xmss/rust/multisig-glue/src/lib.rs b/xmss/rust/multisig-glue/src/lib.rs index 332b420d..e41a7854 100644 --- a/xmss/rust/multisig-glue/src/lib.rs +++ b/xmss/rust/multisig-glue/src/lib.rs @@ -1,12 +1,8 @@ -use backend::symmetric::Permutation; -use backend::{default_koalabear_poseidon1_16, KoalaBear, PrimeField32}; -use lean_multisig::{ - aggregate_single_message_signatures, merge_single_message_aggregates, setup_prover, - setup_prover_without_arena, setup_verifier, verify_multi_message_aggregate, - verify_single_message_aggregate, MultiMessageAggregateSignature, - SingleMessageAggregateSignature, XmssPublicKey, XmssSignature, +use leanvm::xmss::{Epoch, Message, XmssPublicKey, XmssSignature, MESSAGE_LEN}; +use leanvm::{ + aggregate, setup_prover, setup_prover_without_arena, setup_verifier, AggregationError, + ClaimSelection, EthereumProof, SignatureClaims, XmssClaimGroup, }; -use rec_aggregation::split_multi_message_aggregate_by_message; use std::panic::AssertUnwindSafe; use std::slice; use std::sync::OnceLock; @@ -21,8 +17,6 @@ pub struct Signature { pub inner: XmssSignature, } -const MESSAGE_LEN: usize = 32; - static PROVER_READY: OnceLock = OnceLock::new(); static VERIFIER_READY: OnceLock = OnceLock::new(); @@ -32,17 +26,11 @@ macro_rules! ffi_guard { }; } -// setup_prover enables a process-wide arena allocator and warms the prover. -// Its single shared region means two proofs must never be generated -// concurrently; the Go-side proving.Gate serializes all aggregate/merge/split -// work to one at a time, so that invariant holds. Verification does not use the -// arena and stays safe to run concurrently. -// -// The arena is faster but never returns pages to the OS, so RSS ratchets to the -// allocation high-water mark and stays there. xmss_setup_prover_without_arena -// warms the same prover on the system allocator instead: slower, but each proof's -// scratch is freed, keeping a long-lived node's memory bounded. Only one of the -// two is ever called (shared readiness latch), chosen once at startup. +// leanVM allows one `aggregate` call at a time per process, and setup_prover's arena is a +// single shared region. The Go-side proving.Gate serializes all aggregate, merge and split +// work, so that holds. Verification is not affected and stays safe to run concurrently. +// Only one of the two prover setups is ever called (shared readiness latch), chosen once at +// startup. #[no_mangle] pub extern "C" fn xmss_setup_prover() -> i32 { let ready = PROVER_READY.get_or_init(|| std::panic::catch_unwind(setup_prover).is_ok()); @@ -91,11 +79,21 @@ unsafe fn write_out(src: &[u8], out: *mut u8, cap: usize, written: *mut usize) - 0 } +unsafe fn read_message(ptr: *const u8) -> Option { + if ptr.is_null() { + return None; + } + slice::from_raw_parts(ptr, MESSAGE_LEN).try_into().ok() +} + unsafe fn collect_pubkeys( ptrs: *const *const PublicKey, count: usize, ) -> Option> { - if count > 0 && ptrs.is_null() { + if count == 0 { + return Some(Vec::new()); + } + if ptrs.is_null() { return None; } let mut keys = Vec::with_capacity(count); @@ -108,24 +106,167 @@ unsafe fn collect_pubkeys( Some(keys) } -unsafe fn collect_key_groups( - flat: *const *const PublicKey, - counts: *const usize, - group_count: usize, -) -> Option>> { - if group_count == 0 || flat.is_null() || counts.is_null() { +/// Reads `count` claim groups: group `i` holds `pubkey_counts[i]` keys from the flattened +/// `pubkeys`, which signed the `i`th 32-byte message in `message_hashes` at `slots[i]`. +unsafe fn collect_groups( + pubkeys: *const *const PublicKey, + pubkey_counts: *const usize, + message_hashes: *const u8, + slots: *const u32, + count: usize, +) -> Option> { + if count == 0 + || pubkeys.is_null() + || pubkey_counts.is_null() + || message_hashes.is_null() + || slots.is_null() + { return None; } - let counts = slice::from_raw_parts(counts, group_count); - let mut groups = Vec::with_capacity(group_count); + let counts = slice::from_raw_parts(pubkey_counts, count); + let slots = slice::from_raw_parts(slots, count); + let mut groups = Vec::with_capacity(count); let mut offset = 0usize; - for &count in counts { - groups.push(collect_pubkeys(flat.add(offset), count)?); - offset = offset.checked_add(count)?; + for i in 0..count { + groups.push(XmssClaimGroup { + epoch: slots[i], + message: read_message(message_hashes.add(i.checked_mul(MESSAGE_LEN)?))?, + keys: collect_pubkeys(pubkeys.add(offset), counts[i])?, + }); + offset = offset.checked_add(counts[i])?; } Some(groups) } +/// Reads `count` raw signatures: signature `i` is by `keys[i]` over the `i`th 32-byte message +/// in `message_hashes`, at `slots[i]`. +unsafe fn collect_raw( + keys: *const *const PublicKey, + signatures: *const *const Signature, + message_hashes: *const u8, + slots: *const u32, + count: usize, +) -> Option> { + if count == 0 { + return Some(Vec::new()); + } + if keys.is_null() || signatures.is_null() || message_hashes.is_null() || slots.is_null() { + return None; + } + let keys = slice::from_raw_parts(keys, count); + let signatures = slice::from_raw_parts(signatures, count); + let slots = slice::from_raw_parts(slots, count); + let mut raw = Vec::with_capacity(count); + for i in 0..count { + if keys[i].is_null() || signatures[i].is_null() { + return None; + } + raw.push(( + (*keys[i]).inner.clone(), + slots[i], + read_message(message_hashes.add(i.checked_mul(MESSAGE_LEN)?))?, + (*signatures[i]).inner.clone(), + )); + } + Some(raw) +} + +/// The signer set a proof over `groups` is bound to. leanVM requires keys strictly sorted +/// within a group and groups strictly increasing by `(epoch, message)`, and the prover and +/// every verifier must derive the identical set, so it is built here in one place: groups +/// sharing an epoch and message are merged, keys are sorted and deduplicated. An epoch +/// signed at under several messages is one group per message. +fn signature_claims(mut groups: Vec) -> SignatureClaims { + groups.sort_by(|a, b| (a.epoch, a.message).cmp(&(b.epoch, b.message))); + let mut merged: Vec = Vec::with_capacity(groups.len()); + for group in groups { + match merged.last_mut() { + Some(last) if (last.epoch, last.message) == (group.epoch, group.message) => { + last.keys.extend(group.keys); + } + _ => merged.push(group), + } + } + for group in &mut merged { + group.keys.sort(); + group.keys.dedup(); + } + SignatureClaims { + xmss: merged, + sphincs: Vec::new(), + } +} + +fn single_group(epoch: Epoch, message: Message, keys: Vec) -> SignatureClaims { + signature_claims(vec![XmssClaimGroup { + epoch, + message, + keys, + }]) +} + +/// Decodes proof bytes against the claims the caller expects them to prove. The bytes carry +/// no claims of their own, so a proof decoded against the wrong claims fails verification. +unsafe fn decode_proof( + proof: *const u8, + proof_len: usize, + claims: SignatureClaims, +) -> Option { + if proof.is_null() || proof_len == 0 { + return None; + } + EthereumProof::from_bytes_without_pubkeys(slice::from_raw_parts(proof, proof_len), claims).ok() +} + +/// Why a proving call failed, returned as its status so a failure names its cause. The +/// Go side maps each code to a message (xmss/ffi.go). -1 is an invalid argument and -2 an +/// output buffer too small. +#[derive(Clone, Copy)] +#[repr(i32)] +enum Failure { + UndecodableInput = -3, + SplitTargetMissing = -5, + Panicked = -6, + InvalidChild = -10, + MalformedRawSignature = -11, + TooManyEpochs = -12, + TooLarge = -13, + ChildOutOfRange = -14, + NotCovered = -15, + Empty = -16, + InvalidRate = -17, + InvalidBlob = -18, +} + +impl From for Failure { + fn from(err: AggregationError) -> Self { + match err { + AggregationError::InvalidChild(_) => Failure::InvalidChild, + AggregationError::MalformedRawSignature => Failure::MalformedRawSignature, + AggregationError::TooManyEpochs => Failure::TooManyEpochs, + AggregationError::TooLarge => Failure::TooLarge, + AggregationError::ChildOutOfRange { .. } => Failure::ChildOutOfRange, + AggregationError::NotCovered => Failure::NotCovered, + AggregationError::Empty => Failure::Empty, + AggregationError::InvalidRate { .. } => Failure::InvalidRate, + AggregationError::InvalidBlobSize { .. } | AggregationError::BlobNotCovered => { + Failure::InvalidBlob + } + } + } +} + +/// XMSS-only `aggregate`. It panics on an invalid raw signature, so every caller runs it +/// inside `ffi_guard`. +fn prove( + children: &[EthereumProof], + raw: Vec<(XmssPublicKey, Epoch, Message, XmssSignature)>, + declare: Option>, + log_inv_rate: usize, +) -> Result { + aggregate(children, raw, Vec::new(), &[], declare, log_inv_rate).map_err(Failure::from) +} + #[no_mangle] pub unsafe extern "C" fn xmss_aggregate_type_1( raw_pub_keys: *const *const PublicKey, @@ -143,9 +284,8 @@ pub unsafe extern "C" fn xmss_aggregate_type_1( cap: usize, written: *mut usize, ) -> i32 { - ffi_guard!(-1, { - if message_hash.is_null() - || written.is_null() + ffi_guard!(Failure::Panicked as i32, { + if written.is_null() || (num_raw > 0 && (raw_pub_keys.is_null() || raw_signatures.is_null())) || (num_children > 0 && (child_all_pub_keys.is_null() @@ -155,11 +295,9 @@ pub unsafe extern "C" fn xmss_aggregate_type_1( { return -1; } - let message: [u8; MESSAGE_LEN] = - match slice::from_raw_parts(message_hash, MESSAGE_LEN).try_into() { - Ok(message) => message, - Err(_) => return -1, - }; + let Some(message) = read_message(message_hash) else { + return -1; + }; let mut raw = Vec::with_capacity(num_raw); if num_raw > 0 { @@ -169,7 +307,12 @@ pub unsafe extern "C" fn xmss_aggregate_type_1( if keys[i].is_null() || signatures[i].is_null() { return -1; } - raw.push(((*keys[i]).inner.clone(), (*signatures[i]).inner.clone())); + raw.push(( + (*keys[i]).inner.clone(), + slot, + message, + (*signatures[i]).inner.clone(), + )); } } @@ -180,32 +323,25 @@ pub unsafe extern "C" fn xmss_aggregate_type_1( let lengths = slice::from_raw_parts(child_proof_lens, num_children); let mut offset = 0usize; for i in 0..num_children { - let keys = match collect_pubkeys(child_all_pub_keys.add(offset), counts[i]) { - Some(keys) => keys, - None => return -1, - }; - offset = match offset.checked_add(counts[i]) { - Some(offset) => offset, - None => return -1, + let Some(keys) = collect_pubkeys(child_all_pub_keys.add(offset), counts[i]) else { + return -1; }; - if proofs[i].is_null() || lengths[i] == 0 { + let Some(next) = offset.checked_add(counts[i]) else { return -1; - } - let proof = slice::from_raw_parts(proofs[i], lengths[i]); - match SingleMessageAggregateSignature::from_bytes_without_pubkeys(proof, keys) { + }; + offset = next; + let claims = single_group(slot, message, keys); + match decode_proof(proofs[i], lengths[i], claims) { Some(proof) => children.push(proof), - None => return -1, + None => return Failure::UndecodableInput as i32, } } } - let proof = match std::panic::catch_unwind(AssertUnwindSafe(|| { - aggregate_single_message_signatures(&children, raw, message, slot, log_inv_rate) - })) { - Ok(Ok(proof)) => proof, - _ => return -1, - }; - write_out(&proof.to_bytes_without_pubkeys(), out, cap, written) + match prove(&children, raw, None, log_inv_rate) { + Ok(proof) => write_out(&proof.to_bytes_without_pubkeys(), out, cap, written), + Err(failure) => failure as i32, + } }) } @@ -219,89 +355,90 @@ pub unsafe extern "C" fn xmss_verify_type_1( proof_len: usize, ) -> bool { ffi_guard!(false, { - if message_hash.is_null() || proof.is_null() || proof_len == 0 { + let Some(message) = read_message(message_hash) else { return false; - } - let message: [u8; MESSAGE_LEN] = - match slice::from_raw_parts(message_hash, MESSAGE_LEN).try_into() { - Ok(message) => message, - Err(_) => return false, - }; - let keys = match collect_pubkeys(public_keys, num_keys) { - Some(keys) => keys, - None => return false, - }; - let proof = match SingleMessageAggregateSignature::from_bytes_without_pubkeys( - slice::from_raw_parts(proof, proof_len), - keys, - ) { - Some(proof) => proof, - None => return false, }; - if proof.info.core.message != message || proof.info.core.slot != slot { + let Some(claims) = + collect_pubkeys(public_keys, num_keys).map(|keys| single_group(slot, message, keys)) + else { return false; - } - verify_single_message_aggregate(&proof).is_ok() + }; + decode_proof(proof, proof_len, claims).is_some_and(|proof| proof.verify().is_ok()) }) } +/// Merges Type-1 proofs, and raw signatures folded in as they are, into one Type-2. Proving +/// a raw signature on its own first and merging that proof would cost a whole extra proof. #[no_mangle] pub unsafe extern "C" fn xmss_merge_type_1_to_type_2( proof_ptrs: *const *const u8, proof_lens: *const usize, pubkeys: *const *const PublicKey, pubkey_counts: *const usize, + message_hashes: *const u8, + message_slots: *const u32, count: usize, + raw_pub_keys: *const *const PublicKey, + raw_signatures: *const *const Signature, + raw_message_hashes: *const u8, + raw_slots: *const u32, + num_raw: usize, log_inv_rate: usize, out: *mut u8, cap: usize, written: *mut usize, ) -> i32 { - ffi_guard!(-1, { - if count == 0 - || proof_ptrs.is_null() - || proof_lens.is_null() - || pubkeys.is_null() - || pubkey_counts.is_null() - || written.is_null() - { + ffi_guard!(Failure::Panicked as i32, { + if written.is_null() { return -1; } - let proof_ptrs = slice::from_raw_parts(proof_ptrs, count); - let proof_lens = slice::from_raw_parts(proof_lens, count); - let groups = match collect_key_groups(pubkeys, pubkey_counts, count) { - Some(groups) => groups, - None => return -1, - }; - let mut proofs = Vec::with_capacity(count); - for i in 0..count { - if proof_ptrs[i].is_null() || proof_lens[i] == 0 { + let mut children = Vec::with_capacity(count); + if count > 0 { + if proof_ptrs.is_null() || proof_lens.is_null() { return -1; } - match SingleMessageAggregateSignature::from_bytes_without_pubkeys( - slice::from_raw_parts(proof_ptrs[i], proof_lens[i]), - groups[i].clone(), - ) { - Some(proof) => proofs.push(proof), - None => return -1, + let Some(groups) = + collect_groups(pubkeys, pubkey_counts, message_hashes, message_slots, count) + else { + return -1; + }; + let proof_ptrs = slice::from_raw_parts(proof_ptrs, count); + let proof_lens = slice::from_raw_parts(proof_lens, count); + for (i, group) in groups.into_iter().enumerate() { + let claims = signature_claims(vec![group]); + match decode_proof(proof_ptrs[i], proof_lens[i], claims) { + Some(proof) => children.push(proof), + None => return Failure::UndecodableInput as i32, + } } } - let proof = match std::panic::catch_unwind(AssertUnwindSafe(|| { - merge_single_message_aggregates(proofs, log_inv_rate) - })) { - Ok(Ok(proof)) => proof, - _ => return -1, + let Some(raw) = collect_raw( + raw_pub_keys, + raw_signatures, + raw_message_hashes, + raw_slots, + num_raw, + ) else { + return -1; }; - write_out(&proof.to_bytes_without_pubkeys(), out, cap, written) + match prove(&children, raw, None, log_inv_rate) { + Ok(proof) => write_out(&proof.to_bytes_without_pubkeys(), out, cap, written), + Err(failure) => failure as i32, + } }) } +/// Re-proves the one group of a Type-2 proof that signed `target_message` as a standalone +/// Type-1. leanVM has no cheaper split: the Type-2 is a child of a new proof that declares +/// only the kept group. #[no_mangle] pub unsafe extern "C" fn xmss_split_type_2_by_message( proof: *const u8, proof_len: usize, pubkeys: *const *const PublicKey, pubkey_counts: *const usize, + message_hashes: *const u8, + message_slots: *const u32, count: usize, target_message: *const u8, log_inv_rate: usize, @@ -309,33 +446,38 @@ pub unsafe extern "C" fn xmss_split_type_2_by_message( cap: usize, written: *mut usize, ) -> i32 { - ffi_guard!(-1, { - if proof.is_null() || proof_len == 0 || target_message.is_null() || written.is_null() { + ffi_guard!(Failure::Panicked as i32, { + if written.is_null() { return -1; } - let groups = match collect_key_groups(pubkeys, pubkey_counts, count) { - Some(groups) => groups, - None => return -1, + let Some(target) = read_message(target_message) else { + return -1; }; - let proof = match MultiMessageAggregateSignature::from_bytes_without_pubkeys( - slice::from_raw_parts(proof, proof_len), - groups, - ) { - Some(proof) => proof, - None => return -1, + let Some(groups) = + collect_groups(pubkeys, pubkey_counts, message_hashes, message_slots, count) + else { + return -1; }; - let target: [u8; MESSAGE_LEN] = - match slice::from_raw_parts(target_message, MESSAGE_LEN).try_into() { - Ok(target) => target, - Err(_) => return -1, - }; - let proof = match std::panic::catch_unwind(AssertUnwindSafe(|| { - split_multi_message_aggregate_by_message(proof, target, log_inv_rate) - })) { - Ok(Ok(proof)) => proof, - _ => return -1, + let claims = signature_claims(groups); + let mut kept = claims.xmss.iter().filter(|group| group.message == target); + let (Some(group), None) = (kept.next(), kept.next()) else { + return Failure::SplitTargetMissing as i32; + }; + let kept = SignatureClaims { + xmss: vec![group.clone()], + sphincs: Vec::new(), }; - write_out(&proof.to_bytes_without_pubkeys(), out, cap, written) + let Some(type_2) = decode_proof(proof, proof_len, claims) else { + return Failure::UndecodableInput as i32; + }; + let declare = ClaimSelection { + signatures: &kept, + da_commitments: &[], + }; + match prove(&[type_2], Vec::new(), Some(declare), log_inv_rate) { + Ok(proof) => write_out(&proof.to_bytes_without_pubkeys(), out, cap, written), + Err(failure) => failure as i32, + } }) } @@ -350,65 +492,12 @@ pub unsafe extern "C" fn xmss_verify_type_2( message_slots: *const u32, ) -> bool { ffi_guard!(false, { - if proof.is_null() || proof_len == 0 || message_hashes.is_null() || message_slots.is_null() - { + let Some(claims) = + collect_groups(pubkeys, pubkey_counts, message_hashes, message_slots, count) + .map(signature_claims) + else { return false; - } - let groups = match collect_key_groups(pubkeys, pubkey_counts, count) { - Some(groups) => groups, - None => return false, - }; - let proof = match MultiMessageAggregateSignature::from_bytes_without_pubkeys( - slice::from_raw_parts(proof, proof_len), - groups, - ) { - Some(proof) => proof, - None => return false, }; - if proof.info.len() != count { - return false; - } - let hashes = slice::from_raw_parts(message_hashes, count * MESSAGE_LEN); - let slots = slice::from_raw_parts(message_slots, count); - for i in 0..count { - let mut expected = [0; MESSAGE_LEN]; - expected.copy_from_slice(&hashes[i * MESSAGE_LEN..(i + 1) * MESSAGE_LEN]); - if proof.info[i].core.message != expected || proof.info[i].core.slot != slots[i] { - return false; - } - } - verify_multi_message_aggregate(&proof).is_ok() - }) -} - -// Poseidon permutation over KoalaBear, exposed for spec test vectors. The state -// is the canonical u32 field representation, permuted in place. This is the same -// instance the XMSS stack hashes with, so it matches the spec. Returns -1 on a -// null pointer or wrong width. -#[no_mangle] -pub unsafe extern "C" fn poseidon_permute_kb16(state: *mut u32, len: usize) -> i32 { - ffi_guard!(-1, { - if state.is_null() || len != 16 { - return -1; - } - let raw = slice::from_raw_parts_mut(state, 16); - let mut input = [0u32; 16]; - input.copy_from_slice(raw); - let mut fe = KoalaBear::new_array(input); - default_koalabear_poseidon1_16().permute_mut(&mut fe); - for (dst, x) in raw.iter_mut().zip(fe.iter()) { - *dst = x.as_canonical_u32(); - } - 0 + decode_proof(proof, proof_len, claims).is_some_and(|proof| proof.verify().is_ok()) }) } - -// SPIKE FOLLOW-UP: leanVM's internalized XMSS dropped the width-24 Poseidon permutation -// (the scheme now hashes only with width-16 `poseidon16_compress`), so there is no upstream -// constructor to back this. Returns -1 (unsupported) rather than silently mis-permuting. -// The only consumer is the `poseidon_permutation` spec-vector test; revisit when the leanSpec -// pin bumps to the new scheme — the vector set is expected to drop width-24 too. -#[no_mangle] -pub unsafe extern "C" fn poseidon_permute_kb24(_state: *mut u32, _len: usize) -> i32 { - -1 -} diff --git a/xmss/type2_test.go b/xmss/type2_test.go index c88a4a41..7445b272 100644 --- a/xmss/type2_test.go +++ b/xmss/type2_test.go @@ -1,6 +1,9 @@ package xmss -import "testing" +import ( + "fmt" + "testing" +) func TestType2Roundtrip(t *testing.T) { key, err := GenerateKeyPair("type-2-roundtrip", 0, 1<<10) @@ -41,10 +44,10 @@ func TestType2Roundtrip(t *testing.T) { if err != nil { t.Fatal(err) } - inputs = append(inputs, Type1Input{Pubkeys: []CPubKey{pubkey}, Proof: proof}) + inputs = append(inputs, Type1Input{Pubkeys: []CPubKey{pubkey}, Proof: proof, Message: message, Slot: uint32(slot)}) } - proof, err := MergeType1Proofs(inputs) + proof, err := MergeType1Proofs(inputs, nil) if err != nil { t.Fatal(err) } @@ -56,7 +59,7 @@ func TestType2Roundtrip(t *testing.T) { if err := VerifyType2Proof(proof, groups, bindings); err != nil { t.Fatal(err) } - recovered, err := SplitType2Proof(proof, groups, messages[0]) + recovered, err := SplitType2Proof(proof, groups, bindings, messages[0]) if err != nil { t.Fatal(err) } @@ -64,3 +67,130 @@ func TestType2Roundtrip(t *testing.T) { t.Fatal(err) } } + +// A claim group is an (epoch, message) pair, so two validators that signed different +// messages at the same slot are two groups of one proof. Validators disagreeing within a +// slot is ordinary, and the merge has to carry both rather than lose the block. +func TestType2MergesTwoMessagesAtOneSlot(t *testing.T) { + const slot = 5 + var inputs []Type1Input + var bindings []MessageBinding + var groups [][]CPubKey + for i, seed := range []string{"same-slot-a", "same-slot-b"} { + key, err := GenerateKeyPair(seed, 0, 1<<10) + if err != nil { + t.Fatal(err) + } + defer key.Close() + pubkeyBytes, err := key.PublicKeyBytes() + if err != nil { + t.Fatal(err) + } + pubkey, err := ParsePublicKey(pubkeyBytes) + if err != nil { + t.Fatal(err) + } + defer FreePublicKey(pubkey) + + var message [32]byte + message[0] = byte(i + 1) + raw, err := key.Sign(slot, message) + if err != nil { + t.Fatal(err) + } + signature, err := ParseSignature(raw[:]) + if err != nil { + t.Fatal(err) + } + proof, err := AggregateSignatures([]CPubKey{pubkey}, []CSig{signature}, message, slot) + FreeSignature(signature) + if err != nil { + t.Fatal(err) + } + inputs = append(inputs, Type1Input{Pubkeys: []CPubKey{pubkey}, Proof: proof, Message: message, Slot: slot}) + bindings = append(bindings, MessageBinding{Message: message, Slot: slot}) + groups = append(groups, []CPubKey{pubkey}) + } + + proof, err := MergeType1Proofs(inputs, nil) + if err != nil { + t.Fatal(err) + } + if err := VerifyType2Proof(proof, groups, bindings); err != nil { + t.Fatal(err) + } + // Each group is bound to its own message: the proof does not verify against the + // two keys with their messages exchanged. + swapped := []MessageBinding{ + {Message: bindings[1].Message, Slot: slot}, + {Message: bindings[0].Message, Slot: slot}, + } + if VerifyType2Proof(proof, groups, swapped) == nil { + t.Fatal("verified against exchanged messages") + } +} + +// A raw signature merged beside Type-1 proofs is bound to its own message and slot, +// exactly as a proof of it would be. +func TestType2MergesRawSignature(t *testing.T) { + keys := make([]*ValidatorKeyPair, 3) + pubkeys := make([]CPubKey, 3) + for i := range keys { + key, err := GenerateKeyPair(fmt.Sprintf("raw-merge-%d", i), 0, 1<<10) + if err != nil { + t.Fatal(err) + } + defer key.Close() + pubkeyBytes, err := key.PublicKeyBytes() + if err != nil { + t.Fatal(err) + } + pubkey, err := ParsePublicKey(pubkeyBytes) + if err != nil { + t.Fatal(err) + } + defer FreePublicKey(pubkey) + keys[i], pubkeys[i] = key, pubkey + } + sign := func(i int, slot uint32, message [32]byte) CSig { + raw, err := keys[i].Sign(slot, message) + if err != nil { + t.Fatal(err) + } + sig, err := ParseSignature(raw[:]) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { FreeSignature(sig) }) + return sig + } + + var bindings []MessageBinding + var inputs []Type1Input + for i, slot := range []uint32{1, 2} { + var message [32]byte + message[0] = byte(slot) + proof, err := AggregateSignatures([]CPubKey{pubkeys[i]}, []CSig{sign(i, slot, message)}, message, slot) + if err != nil { + t.Fatal(err) + } + inputs = append(inputs, Type1Input{Pubkeys: []CPubKey{pubkeys[i]}, Proof: proof, Message: message, Slot: slot}) + bindings = append(bindings, MessageBinding{Message: message, Slot: slot}) + } + var blockRoot [32]byte + blockRoot[0] = 0xB + proposer := RawSignature{Pubkey: pubkeys[2], Signature: sign(2, 3, blockRoot), Message: blockRoot, Slot: 3} + + proof, err := MergeType1Proofs(inputs, []RawSignature{proposer}) + if err != nil { + t.Fatal(err) + } + groups := [][]CPubKey{{pubkeys[0]}, {pubkeys[1]}, {pubkeys[2]}} + if err := VerifyType2Proof(proof, groups, append(bindings, MessageBinding{Message: blockRoot, Slot: 3})); err != nil { + t.Fatal(err) + } + var otherRoot [32]byte + if VerifyType2Proof(proof, groups, append(bindings, MessageBinding{Message: otherRoot, Slot: 3})) == nil { + t.Fatal("verified against a different proposer message") + } +}