1+ use std:: sync:: Mutex ;
2+
13use backend:: * ;
24use rand:: { CryptoRng , RngExt , SeedableRng , rngs:: StdRng } ;
35use serde:: { Deserialize , Serialize } ;
46use sha3:: { Digest as Sha3Digest , Keccak256 } ;
57
68use crate :: * ;
79
10+ /// Memory-optimized secret key for a range of R = slot_end - slot_start + 1 slots: O(sqrt(R) +
11+ /// LOG_LIFETIME) instead of O(R). Stores the top tree (in-range band plus a thin spine) and one
12+ /// cached bottom subtree, cut at split_level = log2(R)/2. Out-of-range nodes are deterministic
13+ /// gen_random_node fillers; see `xmss_small_memory.tex` for the picture.
814#[ derive( Debug ) ]
915pub struct XmssSecretKey {
1016 pub ( crate ) slot_start : u32 , // inclusive
1117 pub ( crate ) slot_end : u32 , // inclusive
1218 pub ( crate ) public_param : PublicParam ,
1319 pub ( crate ) seed : [ u8 ; 32 ] ,
14- // At level l, stored indices go from (slot_start >> l) to (slot_end >> l).
15- pub ( crate ) merkle_tree : Vec < Vec < Digest > > ,
20+ pub ( crate ) split_level : usize , // bottom-subtree height (2^split_level leaves each)
21+ // top[l - split_level] = level-l nodes for indices [slot_start >> l, slot_end >> l]
22+ pub ( crate ) top : Vec < Vec < Digest > > ,
23+ pub ( crate ) cache : Mutex < Option < BottomSubtree > > ,
24+ }
25+
26+ /// Bottom subtree covering the last-signed slot; its leaf range is derived from `subtree_index`.
27+ #[ derive( Debug ) ]
28+ pub ( crate ) struct BottomSubtree {
29+ subtree_index : u64 , // = slot >> split_level
30+ layers : Vec < Vec < Digest > > ,
1631}
1732
1833#[ derive( Debug , Clone , Serialize , Deserialize , PartialEq , Eq , Hash , PartialOrd , Ord ) ]
@@ -81,73 +96,137 @@ fn fill<T: Send>(sequential: bool, data: &mut [T], f: impl Fn(usize, &mut T) + S
8196 }
8297}
8398
84- pub fn xmss_key_gen (
85- seed : [ u8 ; 32 ] ,
86- slot_start : u32 ,
87- slot_end : u32 ,
88- sequential : bool ,
89- ) -> Result < ( XmssSecretKey , XmssPublicKey ) , XmssKeyGenError > {
90- if slot_start > slot_end || slot_end as u64 >= ( 1 << LOG_LIFETIME ) {
91- return Err ( XmssKeyGenError :: InvalidRange ) ;
92- }
93- let public_param: PublicParam = gen_public_param ( & seed) ;
94- // Level 0: WOTS leaf hashes for slots in [slot_start, slot_end]
95- let n_leaves = ( slot_end - slot_start + 1 ) as usize ;
96- let mut leaves: Vec < Digest > = unsafe { uninitialized_vec ( n_leaves) } ;
97- fill ( sequential, & mut leaves, |i, out| {
98- let slot = slot_start + i as u32 ;
99- let wots = gen_wots_secret_key ( & seed, slot, public_param) ;
100- * out = wots. public_key ( ) . hash ( public_param, slot) ;
99+ /// Level-0 layer: WOTS public-key hashes for the in-range leaves `[lo, hi]`.
100+ fn leaf_layer ( seed : & [ u8 ; 32 ] , public_param : & PublicParam , lo : u64 , hi : u64 , sequential : bool ) -> Vec < Digest > {
101+ let mut leaves: Vec < Digest > = unsafe { uninitialized_vec ( ( hi - lo + 1 ) as usize ) } ;
102+ fill ( sequential, & mut leaves, |k, out| {
103+ let slot = ( lo + k as u64 ) as u32 ;
104+ let wots = gen_wots_secret_key ( seed, slot, * public_param) ;
105+ * out = wots. public_key ( ) . hash ( * public_param, slot) ;
101106 } ) ;
102- let mut merkle_tree = vec ! [ leaves] ;
103- // Build levels 1..=LOG_LIFETIME.
104- // At level l, we store nodes with index in [(slot_start >> l), (slot_end >> l)].
105- // Children outside [slot_start, slot_end]'s subtree are replaced by gen_random_node.
106- for level in 1 ..=LOG_LIFETIME {
107- let base: u64 = ( slot_start as u64 ) >> level;
108- let top: u64 = ( slot_end as u64 ) >> level;
109- let prev_base: u64 = ( slot_start as u64 ) >> ( level - 1 ) ;
110- let prev_top: u64 = ( slot_end as u64 ) >> ( level - 1 ) ;
107+ leaves
108+ }
109+
110+ /// Build levels `(from_level+1)..=to_level` onto `layers`; out-of-range children use `gen_random_node`.
111+ #[ allow( clippy:: too_many_arguments) ]
112+ fn build_up (
113+ seed : & [ u8 ; 32 ] ,
114+ public_param : & PublicParam ,
115+ layers : & mut Vec < Vec < Digest > > ,
116+ lo : u64 ,
117+ hi : u64 ,
118+ from_level : usize ,
119+ to_level : usize ,
120+ sequential : bool ,
121+ ) {
122+ for level in ( from_level + 1 ) ..=to_level {
123+ let base = lo >> level;
124+ let top = hi >> level;
125+ let prev_base = lo >> ( level - 1 ) ;
126+ let prev_top = hi >> ( level - 1 ) ;
111127 let nodes: Vec < Digest > = {
112- let prev = & merkle_tree[ level - 1 ] ;
113- let n_nodes = ( top - base + 1 ) as usize ;
114- let mut nodes: Vec < Digest > = unsafe { uninitialized_vec ( n_nodes) } ;
128+ let prev = layers. last ( ) . unwrap ( ) ;
129+ let mut nodes: Vec < Digest > = unsafe { uninitialized_vec ( ( top - base + 1 ) as usize ) } ;
115130 fill ( sequential, & mut nodes, |k, out| {
116131 let i = base + k as u64 ;
117132 let left_idx = 2 * i;
118133 let right_idx = 2 * i + 1 ;
119134 let left = if left_idx >= prev_base && left_idx <= prev_top {
120135 prev[ ( left_idx - prev_base) as usize ]
121136 } else {
122- gen_random_node ( & seed, level - 1 , left_idx)
137+ gen_random_node ( seed, level - 1 , left_idx)
123138 } ;
124139 let right = if right_idx >= prev_base && right_idx <= prev_top {
125140 prev[ ( right_idx - prev_base) as usize ]
126141 } else {
127- gen_random_node ( & seed, level - 1 , right_idx)
142+ gen_random_node ( seed, level - 1 , right_idx)
128143 } ;
129144 let merkle_data = build_merkle_data (
130145 make_tweak ( TWEAK_TYPE_MERKLE , level, i as u32 ) ,
131- & public_param,
146+ public_param,
132147 & left,
133148 & right,
134149 ) ;
135150 * out = poseidon8_compress ( merkle_data) [ ..XMSS_DIGEST_LEN ] . try_into ( ) . unwrap ( ) ;
136151 } ) ;
137152 nodes
138153 } ;
139- merkle_tree . push ( nodes) ;
154+ layers . push ( nodes) ;
140155 }
156+ }
157+
158+ /// In-range leaf bounds of the bottom subtree with the given index.
159+ fn subtree_bounds ( slot_start : u64 , slot_end : u64 , split_level : usize , subtree_index : u64 ) -> ( u64 , u64 ) {
160+ (
161+ slot_start. max ( subtree_index << split_level) ,
162+ slot_end. min ( ( ( subtree_index + 1 ) << split_level) - 1 ) ,
163+ )
164+ }
165+
166+ /// Build merkle layers `0..=to_level` for the in-range leaves `[lo, hi]`.
167+ fn build_subtree_layers (
168+ seed : & [ u8 ; 32 ] ,
169+ public_param : & PublicParam ,
170+ lo : u64 ,
171+ hi : u64 ,
172+ to_level : usize ,
173+ sequential : bool ,
174+ ) -> Vec < Vec < Digest > > {
175+ let mut layers = vec ! [ leaf_layer( seed, public_param, lo, hi, sequential) ] ;
176+ build_up ( seed, public_param, & mut layers, lo, hi, 0 , to_level, sequential) ;
177+ layers
178+ }
179+
180+ pub fn xmss_key_gen (
181+ seed : [ u8 ; 32 ] ,
182+ slot_start : u32 ,
183+ slot_end : u32 ,
184+ sequential : bool ,
185+ ) -> Result < ( XmssSecretKey , XmssPublicKey ) , XmssKeyGenError > {
186+ if slot_start > slot_end || slot_end as u64 >= ( 1 << LOG_LIFETIME ) {
187+ return Err ( XmssKeyGenError :: InvalidRange ) ;
188+ }
189+ let public_param: PublicParam = gen_public_param ( & seed) ;
190+ let lo = slot_start as u64 ;
191+ let hi = slot_end as u64 ;
192+
193+ // ~sqrt(R) leaves per bottom subtree; always <= LOG_LIFETIME/2 since R <= 2^LOG_LIFETIME.
194+ let split_level = log2_ceil_usize ( ( hi - lo + 1 ) as usize ) . div_ceil ( 2 ) ;
195+
196+ // Roots of each bottom subtree, built one at a time so peak memory stays O(sqrt(R)).
197+ let first_subtree = lo >> split_level;
198+ let last_subtree = hi >> split_level;
199+ let mut root_layer: Vec < Digest > = unsafe { uninitialized_vec ( ( last_subtree - first_subtree + 1 ) as usize ) } ;
200+ fill ( sequential, & mut root_layer, |k, out| {
201+ let ( in_lo, in_hi) = subtree_bounds ( lo, hi, split_level, first_subtree + k as u64 ) ;
202+ * out = build_subtree_layers ( & seed, & public_param, in_lo, in_hi, split_level, true ) [ split_level] [ 0 ] ;
203+ } ) ;
204+
205+ // Top part: levels split_level..=LOG_LIFETIME.
206+ let mut top = vec ! [ root_layer] ;
207+ build_up (
208+ & seed,
209+ & public_param,
210+ & mut top,
211+ lo,
212+ hi,
213+ split_level,
214+ LOG_LIFETIME ,
215+ sequential,
216+ ) ;
217+
141218 let pub_key = XmssPublicKey {
142- merkle_root : merkle_tree . last ( ) . unwrap ( ) [ 0 ] ,
219+ merkle_root : top . last ( ) . unwrap ( ) [ 0 ] ,
143220 public_param,
144221 } ;
145222 let secret_key = XmssSecretKey {
146223 slot_start,
147224 slot_end,
148225 public_param,
149226 seed,
150- merkle_tree,
227+ split_level,
228+ top,
229+ cache : Mutex :: new ( None ) ,
151230 } ;
152231 Ok ( ( secret_key, pub_key) )
153232}
@@ -181,16 +260,18 @@ pub fn xmss_sign_with_randomness(
181260 let wots_signature = wots_secret_key
182261 . sign_with_randomness ( message, slot, & secret_key. public_key ( ) , randomness)
183262 . ok_or ( XmssSignatureError :: InvalidRandomness ) ?;
263+ // Cache the bottom subtree covering `slot` (reused across its 2^split_level slots), then read the path.
264+ let subtree_index = ( slot as u64 ) >> secret_key. split_level ;
265+ let mut cache = secret_key. cache . lock ( ) . unwrap ( ) ;
266+ if cache. as_ref ( ) . is_none_or ( |s| s. subtree_index != subtree_index) {
267+ * cache = Some ( secret_key. build_bottom_subtree ( subtree_index) ) ;
268+ }
269+ let sub = cache. as_ref ( ) . unwrap ( ) ;
184270 let merkle_proof = std:: array:: from_fn ( |level| {
185271 let neighbour_index = ( ( slot as u64 ) >> level) ^ 1 ;
186- let base = ( secret_key. slot_start as u64 ) >> level;
187- let top = ( secret_key. slot_end as u64 ) >> level;
188- if neighbour_index >= base && neighbour_index <= top {
189- secret_key. merkle_tree [ level] [ ( neighbour_index - base) as usize ]
190- } else {
191- gen_random_node ( & secret_key. seed , level, neighbour_index)
192- }
272+ secret_key. merkle_sibling ( level, neighbour_index, sub)
193273 } ) ;
274+ drop ( cache) ;
194275 Ok ( XmssSignature {
195276 wots_signature,
196277 merkle_proof,
@@ -200,10 +281,48 @@ pub fn xmss_sign_with_randomness(
200281impl XmssSecretKey {
201282 pub fn public_key ( & self ) -> XmssPublicKey {
202283 XmssPublicKey {
203- merkle_root : self . merkle_tree . last ( ) . unwrap ( ) [ 0 ] ,
284+ merkle_root : self . top . last ( ) . unwrap ( ) [ 0 ] ,
204285 public_param : self . public_param ,
205286 }
206287 }
288+
289+ /// (Re)build the bottom subtree with the given index.
290+ fn build_bottom_subtree ( & self , subtree_index : u64 ) -> BottomSubtree {
291+ let ( lo, hi) = subtree_bounds (
292+ self . slot_start as u64 ,
293+ self . slot_end as u64 ,
294+ self . split_level ,
295+ subtree_index,
296+ ) ;
297+ let layers = build_subtree_layers ( & self . seed , & self . public_param , lo, hi, self . split_level , true ) ;
298+ BottomSubtree { subtree_index, layers }
299+ }
300+
301+ /// Authentication-path sibling at `level`: from the top part, the cached subtree, or `gen_random_node`.
302+ fn merkle_sibling ( & self , level : usize , neighbour_index : u64 , sub : & BottomSubtree ) -> Digest {
303+ let ( lo, hi, level_base, layers) = if level >= self . split_level {
304+ (
305+ self . slot_start as u64 ,
306+ self . slot_end as u64 ,
307+ self . split_level ,
308+ & self . top ,
309+ )
310+ } else {
311+ let ( lo, hi) = subtree_bounds (
312+ self . slot_start as u64 ,
313+ self . slot_end as u64 ,
314+ self . split_level ,
315+ sub. subtree_index ,
316+ ) ;
317+ ( lo, hi, 0 , & sub. layers )
318+ } ;
319+ let base = lo >> level;
320+ if neighbour_index >= base && neighbour_index <= ( hi >> level) {
321+ layers[ level - level_base] [ ( neighbour_index - base) as usize ]
322+ } else {
323+ gen_random_node ( & self . seed , level, neighbour_index)
324+ }
325+ }
207326}
208327
209328#[ derive( Debug , PartialEq , Eq , Clone , Copy , Hash ) ]
0 commit comments