@@ -116,7 +116,7 @@ pub struct MultiUseSandbox {
116116///
117117/// Returns a list of root page table GPAs to walk. If the list is
118118/// empty, only `root_pt_gpa` is used.
119- pub type PtRootFinder = Box < dyn Fn ( & [ u8 ] , & [ u8 ] , u64 ) -> Vec < u64 > + Send > ;
119+ pub type PtRootFinder = Arc < dyn Fn ( & [ u8 ] , & [ u8 ] , u64 ) -> Vec < u64 > + Send + Sync > ;
120120
121121impl MultiUseSandbox {
122122 fn ensure_usable ( & self ) -> Result < ( ) > {
@@ -157,8 +157,12 @@ impl MultiUseSandbox {
157157 /// Set a callback that discovers page table roots from guest memory.
158158 /// The callback receives (snapshot_mem, scratch_mem, cr3) and returns
159159 /// the list of root GPAs to walk during snapshot creation.
160+ ///
161+ /// In-memory snapshots retain the finder across restore. The finder is not
162+ /// serialized.
160163 pub fn set_pt_root_finder ( & mut self , finder : PtRootFinder ) {
161164 self . pt_root_finder = Some ( finder) ;
165+ self . snapshot = None ;
162166 }
163167
164168 /// Create a `MultiUseSandbox` directly from a [`Snapshot`],
@@ -332,7 +336,8 @@ impl MultiUseSandbox {
332336 } ) ?;
333337 }
334338
335- let sbox = MultiUseSandbox :: from_uninit ( host_funcs, hshm, vm) ;
339+ let mut sbox = MultiUseSandbox :: from_uninit ( host_funcs, hshm, vm) ;
340+ sbox. pt_root_finder = snapshot. pt_root_finder ( ) . cloned ( ) ;
336341 Ok ( sbox)
337342 }
338343
@@ -424,6 +429,7 @@ impl MultiUseSandbox {
424429 msrs,
425430 next_action,
426431 host_functions,
432+ self . pt_root_finder . clone ( ) ,
427433 ) ?;
428434 let snapshot = Arc :: new ( memory_snapshot) ;
429435 self . snapshot = Some ( snapshot. clone ( ) ) ;
@@ -619,7 +625,7 @@ impl MultiUseSandbox {
619625
620626 self . mem_mgr
621627 . request_libc_rng_reseed ( rand:: random :: < u32 > ( ) ) ?;
622- self . pt_root_finder = None ;
628+ self . pt_root_finder = snapshot . pt_root_finder ( ) . cloned ( ) ;
623629
624630 // The restored snapshot is now our most current snapshot
625631 self . snapshot = Some ( snapshot. clone ( ) ) ;
@@ -1189,6 +1195,7 @@ fn warn_on_layout_override(
11891195
11901196#[ cfg( test) ]
11911197mod tests {
1198+ use std:: sync:: atomic:: { AtomicUsize , Ordering } ;
11921199 use std:: sync:: { Arc , Barrier } ;
11931200 use std:: thread;
11941201
@@ -1202,6 +1209,7 @@ mod tests {
12021209 use crate :: mem:: memory_region:: { MemoryRegion , MemoryRegionFlags , MemoryRegionType } ;
12031210 use crate :: mem:: shared_mem:: { ExclusiveSharedMemory , GuestSharedMemory , SharedMemory as _} ;
12041211 use crate :: sandbox:: SandboxConfiguration ;
1212+ use crate :: sandbox:: snapshot:: Snapshot ;
12051213 use crate :: sandbox:: uninitialized:: { GuestBlob , GuestEnvironment } ;
12061214 use crate :: {
12071215 GuestBinary , HyperlightError , MultiUseSandbox , Result , SandboxStatus , UninitializedSandbox ,
@@ -1222,6 +1230,23 @@ mod tests {
12221230 assert ! ( SandboxStatus :: Unrecoverable . is_unrecoverable( ) ) ;
12231231 }
12241232
1233+ trait AmbiguousIfSync < Marker > {
1234+ fn assert_not_sync ( ) { }
1235+ }
1236+
1237+ impl < T : ?Sized > AmbiguousIfSync < ( ) > for T { }
1238+ impl < T : ?Sized + Sync > AmbiguousIfSync < u8 > for T { }
1239+
1240+ #[ test]
1241+ fn snapshot_and_sandbox_thread_safety ( ) {
1242+ fn assert_send < T : Send > ( ) { }
1243+ fn assert_send_sync < T : Send + Sync > ( ) { }
1244+
1245+ assert_send :: < MultiUseSandbox > ( ) ;
1246+ let _ = <MultiUseSandbox as AmbiguousIfSync < _ > >:: assert_not_sync;
1247+ assert_send_sync :: < Snapshot > ( ) ;
1248+ }
1249+
12251250 #[ test]
12261251 fn poison ( ) {
12271252 let mut sbox: MultiUseSandbox = {
@@ -2281,6 +2306,8 @@ mod tests {
22812306 . unwrap ( )
22822307 . evolve ( )
22832308 . unwrap ( ) ;
2309+ let source_finder: crate :: sandbox:: PtRootFinder = Arc :: new ( |_, _, root| vec ! [ root] ) ;
2310+ source. set_pt_root_finder ( source_finder. clone ( ) ) ;
22842311 let mut target =
22852312 UninitializedSandbox :: new ( GuestBinary :: FilePath ( simple_guest_as_pathbuf ( ) ) , None )
22862313 . unwrap ( )
@@ -2289,8 +2316,7 @@ mod tests {
22892316
22902317 assert_eq ! ( source. call:: <i32 >( "StackAllocate" , 256i32 ) . unwrap( ) , 256 ) ;
22912318 assert_eq ! ( target. call:: <i32 >( "AddToStatic" , 17i32 ) . unwrap( ) , 17 ) ;
2292- target. set_pt_root_finder ( Box :: new ( |_, _, root| vec ! [ root] ) ) ;
2293- assert ! ( target. pt_root_finder. is_some( ) ) ;
2319+ target. set_pt_root_finder ( Arc :: new ( |_, _, _| Vec :: new ( ) ) ) ;
22942320
22952321 assert_ne ! (
22962322 source. mem_mgr. layout. code_size( ) ,
@@ -2307,7 +2333,10 @@ mod tests {
23072333
23082334 let snapshot = source. snapshot ( ) . unwrap ( ) ;
23092335 target. restore ( snapshot) . unwrap ( ) ;
2310- assert ! ( target. pt_root_finder. is_none( ) ) ;
2336+ assert ! ( Arc :: ptr_eq(
2337+ target. pt_root_finder. as_ref( ) . unwrap( ) ,
2338+ & source_finder
2339+ ) ) ;
23112340 assert_eq ! ( target. call:: <i32 >( "StackAllocate" , 512i32 ) . unwrap( ) , 512 ) ;
23122341 assert ! ( matches!(
23132342 target. call:: <i32 >( "GetStatic" , ( ) ) ,
@@ -2318,6 +2347,68 @@ mod tests {
23182347 ) ) ;
23192348 }
23202349
2350+ #[ test]
2351+ fn snapshot_restore_clears_absent_pt_root_finder ( ) {
2352+ let path = simple_guest_as_pathbuf ( ) ;
2353+ let mut source = UninitializedSandbox :: new ( GuestBinary :: FilePath ( path) , None )
2354+ . unwrap ( )
2355+ . evolve ( )
2356+ . unwrap ( ) ;
2357+ let snapshot = source. snapshot ( ) . unwrap ( ) ;
2358+ assert ! ( snapshot. pt_root_finder( ) . is_none( ) ) ;
2359+
2360+ let path = simple_guest_as_pathbuf ( ) ;
2361+ let mut target = UninitializedSandbox :: new ( GuestBinary :: FilePath ( path) , None )
2362+ . unwrap ( )
2363+ . evolve ( )
2364+ . unwrap ( ) ;
2365+ target. set_pt_root_finder ( Arc :: new ( |_, _, root| vec ! [ root] ) ) ;
2366+
2367+ target. restore ( snapshot) . unwrap ( ) ;
2368+ assert ! ( target. pt_root_finder. is_none( ) ) ;
2369+ }
2370+
2371+ #[ test]
2372+ fn snapshot_restore_uses_retained_pt_root_finder ( ) {
2373+ let source_calls = Arc :: new ( AtomicUsize :: new ( 0 ) ) ;
2374+ let source_calls_in_finder = source_calls. clone ( ) ;
2375+ let source_finder: crate :: sandbox:: PtRootFinder = Arc :: new ( move |_, _, _| {
2376+ source_calls_in_finder. fetch_add ( 1 , Ordering :: Relaxed ) ;
2377+ Vec :: new ( )
2378+ } ) ;
2379+ let path = simple_guest_as_pathbuf ( ) ;
2380+ let mut source = UninitializedSandbox :: new ( GuestBinary :: FilePath ( path) , None )
2381+ . unwrap ( )
2382+ . evolve ( )
2383+ . unwrap ( ) ;
2384+ source. set_pt_root_finder ( source_finder) ;
2385+ let snapshot = source. snapshot ( ) . unwrap ( ) ;
2386+
2387+ let target_calls = Arc :: new ( AtomicUsize :: new ( 0 ) ) ;
2388+ let target_calls_in_finder = target_calls. clone ( ) ;
2389+ let target_finder: crate :: sandbox:: PtRootFinder = Arc :: new ( move |_, _, root| {
2390+ target_calls_in_finder. fetch_add ( 1 , Ordering :: Relaxed ) ;
2391+ vec ! [ root]
2392+ } ) ;
2393+ let path = simple_guest_as_pathbuf ( ) ;
2394+ let mut target = UninitializedSandbox :: new ( GuestBinary :: FilePath ( path) , None )
2395+ . unwrap ( )
2396+ . evolve ( )
2397+ . unwrap ( ) ;
2398+ target. set_pt_root_finder ( target_finder) ;
2399+ target. restore ( snapshot) . unwrap ( ) ;
2400+
2401+ let source_calls_before = source_calls. load ( Ordering :: Relaxed ) ;
2402+ target. call :: < i32 > ( "GetStatic" , ( ) ) . unwrap ( ) ;
2403+ target. snapshot ( ) . unwrap ( ) ;
2404+
2405+ assert_eq ! (
2406+ source_calls. load( Ordering :: Relaxed ) ,
2407+ source_calls_before + 1
2408+ ) ;
2409+ assert_eq ! ( target_calls. load( Ordering :: Relaxed ) , 0 ) ;
2410+ }
2411+
23212412 #[ test]
23222413 fn snapshot_restore_replaces_c_guest_with_rust_guest ( ) {
23232414 let mut source =
0 commit comments