diff --git a/core/src/abi.rs b/core/src/abi.rs index 4160f00..87c5a61 100644 --- a/core/src/abi.rs +++ b/core/src/abi.rs @@ -4,7 +4,7 @@ use std::ffi::c_void; use std::mem::size_of; use std::panic::{catch_unwind, AssertUnwindSafe}; use std::ptr::{self, NonNull}; -use std::sync::atomic::{AtomicBool, AtomicU64, AtomicUsize, Ordering}; +use std::sync::atomic::{AtomicBool, AtomicU64, AtomicU8, AtomicUsize, Ordering}; use std::sync::mpsc::{self, Receiver, RecvTimeoutError, Sender, SyncSender, TrySendError}; use std::sync::{Arc, Mutex, MutexGuard, OnceLock, TryLockError}; use std::thread::{self, JoinHandle}; @@ -44,6 +44,10 @@ const MAX_CONNECT_BYTES: usize = 1024 * 1024; const CALLBACK_QUEUE_CAPACITY: usize = 64; const CONTROL_QUEUE_CAPACITY: usize = 64; +const TRANSITION_OPEN: u8 = 0; +const TRANSITION_CONNECTED: u8 = 1; +const TRANSITION_CANCELLED: u8 = 2; + static NEXT_ID: AtomicUsize = AtomicUsize::new(1); static HANDLES: OnceLock>>> = OnceLock::new(); @@ -53,6 +57,12 @@ static CANCEL_ADMISSION_HOOK: OnceLock>> = Onc #[cfg(test)] static FINAL_CONNECT_HOOK: OnceLock>> = OnceLock::new(); +#[cfg(test)] +static INITIAL_CONNECT_HOOK: OnceLock>> = OnceLock::new(); + +#[cfg(test)] +static FINAL_PUBLISH_HOOK: OnceLock>> = OnceLock::new(); + #[cfg(test)] struct CancelAdmissionHook { core: usize, @@ -67,6 +77,13 @@ struct FinalConnectHook { resume: Receiver<()>, } +#[cfg(test)] +struct ConnectTransitionHook { + core: usize, + reached: Sender<()>, + resume: Receiver<()>, +} + thread_local! { static CALLBACK_MODE: Cell = const { Cell::new(CallbackMode::None) }; } @@ -228,7 +245,7 @@ struct CoreInner { worker_done: Mutex>, worker: Mutex>>, in_flight: AtomicUsize, - cancelled: AtomicBool, + transition: AtomicU8, cancellation_emitted: AtomicBool, freeing: AtomicBool, destroying: AtomicBool, @@ -309,6 +326,26 @@ fn pause_before_connected(core: *mut CoreHandle) { } } +#[cfg(test)] +fn pause_connect_transition( + hooks: &OnceLock>>, + core: *mut CoreHandle, +) { + let hook = hooks.get_or_init(|| Mutex::new(None)); + let selected = { + let mut hook = lock(hook); + if hook.as_ref().is_some_and(|hook| hook.core == core as usize) { + hook.take() + } else { + None + } + }; + if let Some(hook) = selected { + hook.reached.send(()).expect("announce connect transition"); + hook.resume.recv().expect("resume connect transition"); + } +} + fn status(error: CoreError) -> u32 { match error { CoreError::QueueFull => QUEUE_FULL, @@ -450,8 +487,32 @@ fn apply_cancellation(inner: &CoreInner, state: &mut SessionState) { } } +fn cancellation_admitted(inner: &CoreInner) -> bool { + inner.transition.load(Ordering::Acquire) == TRANSITION_CANCELLED +} + +fn commit_cancellation(inner: &CoreInner) -> u32 { + let mut state = lock(&inner.state); + apply_cancellation(inner, &mut state); + CANCELLED +} + +fn claim_connected(inner: &CoreInner) -> bool { + inner + .transition + .compare_exchange( + TRANSITION_OPEN, + TRANSITION_CONNECTED, + Ordering::AcqRel, + Ordering::Acquire, + ) + .is_ok() +} + fn cancel_inner(inner: &CoreInner) { - inner.cancelled.store(true, Ordering::Release); + inner + .transition + .store(TRANSITION_CANCELLED, Ordering::Release); match inner.state.try_lock() { Ok(mut state) => apply_cancellation(inner, &mut state), Err(TryLockError::Poisoned(error)) => { @@ -622,7 +683,7 @@ unsafe extern "C" fn verse_core_create_v1( worker_done: Mutex::new(done_rx), worker: Mutex::new(Some(worker)), in_flight: AtomicUsize::new(0), - cancelled: AtomicBool::new(false), + transition: AtomicU8::new(TRANSITION_OPEN), cancellation_emitted: AtomicBool::new(false), freeing: AtomicBool::new(false), destroying: AtomicBool::new(false), @@ -676,10 +737,12 @@ unsafe extern "C" fn verse_core_connect_v1( } { let mut state = lock(&handle.state); - if handle.cancelled.load(Ordering::Acquire) { + if cancellation_admitted(&handle) { apply_cancellation(&handle, &mut state); return CANCELLED; } + #[cfg(test)] + pause_connect_transition(&INITIAL_CONNECT_HOOK, core); if state.lifecycle != Lifecycle::Created { return INVALID_STATE; } @@ -696,8 +759,8 @@ unsafe extern "C" fn verse_core_connect_v1( cancel_inner(&handle); return admission; } - if handle.cancelled.load(Ordering::Acquire) { - return CANCELLED; + if cancellation_admitted(&handle) { + return commit_cancellation(&handle); } let tls = call_signer( SignerPurpose::Tls, @@ -711,30 +774,31 @@ unsafe extern "C" fn verse_core_connect_v1( } { let mut state = lock(&handle.state); - if handle.cancelled.load(Ordering::Acquire) || state.lifecycle == Lifecycle::Destroying - { + if cancellation_admitted(&handle) || state.lifecycle == Lifecycle::Destroying { apply_cancellation(&handle, &mut state); return CANCELLED; } #[cfg(test)] pause_before_connected(core); - if handle.cancelled.load(Ordering::Acquire) { + if cancellation_admitted(&handle) { + apply_cancellation(&handle, &mut state); + return CANCELLED; + } + #[cfg(test)] + pause_connect_transition(&FINAL_PUBLISH_HOOK, core); + if !claim_connected(&handle) { apply_cancellation(&handle, &mut state); return CANCELLED; } state.lifecycle = Lifecycle::Connected; enqueue_callback_under_state_lock(&handle, CallbackEvent::State(STATE_CONNECTED)); - if handle.cancelled.load(Ordering::Acquire) { + if cancellation_admitted(&handle) { apply_cancellation(&handle, &mut state); return CANCELLED; } } - if handle.cancelled.load(Ordering::Acquire) - && !handle.cancellation_emitted.load(Ordering::Acquire) - { - let mut state = lock(&handle.state); - apply_cancellation(&handle, &mut state); - return CANCELLED; + if cancellation_admitted(&handle) && !handle.cancellation_emitted.load(Ordering::Acquire) { + return commit_cancellation(&handle); } OK }) @@ -760,7 +824,7 @@ unsafe extern "C" fn verse_core_send_input_v1( Err(error) => return error, }; let mut state = lock(&handle.state); - if handle.cancelled.load(Ordering::Acquire) { + if cancellation_admitted(&handle) { apply_cancellation(&handle, &mut state); return CANCELLED; } @@ -795,7 +859,7 @@ unsafe extern "C" fn verse_core_request_idr_v1(core: *mut CoreHandle) -> u32 { Err(error) => return error, }; let mut state = lock(&handle.state); - if handle.cancelled.load(Ordering::Acquire) { + if cancellation_admitted(&handle) { apply_cancellation(&handle, &mut state); return CANCELLED; } @@ -921,7 +985,9 @@ unsafe extern "C" fn verse_core_destroy_v1(core: *mut CoreHandle, timeout_ms: u3 return BUSY; } inner.callbacks_closed.store(true, Ordering::Release); - inner.cancelled.store(true, Ordering::Release); + inner + .transition + .store(TRANSITION_CANCELLED, Ordering::Release); { if Instant::now() > deadline { return destroy_busy(&inner); @@ -983,6 +1049,7 @@ mod tests { use std::ffi::c_void; use std::mem::size_of; use std::ptr; + use std::sync::atomic::Ordering; use std::sync::mpsc; use std::sync::{Arc, Condvar, Mutex}; use std::thread; @@ -990,9 +1057,9 @@ mod tests { use super::{ ffi_boundary, handles, lock, BytesView, CancelAdmissionHook, Config, ConnectRequest, - CoreHandle, FinalConnectHook, SignFn, StateEvent, StateFn, BUSY, CANCELLED, - CANCEL_ADMISSION_HOOK, FINAL_CONNECT_HOOK, INTERNAL, OK, STATE_CANCELLED, STATE_CONNECTED, - STATE_CONNECTING, + ConnectTransitionHook, CoreHandle, FinalConnectHook, Lifecycle, SignFn, StateEvent, + StateFn, BUSY, CANCELLED, CANCEL_ADMISSION_HOOK, FINAL_CONNECT_HOOK, FINAL_PUBLISH_HOOK, + INITIAL_CONNECT_HOOK, INTERNAL, OK, STATE_CANCELLED, STATE_CONNECTED, STATE_CONNECTING, }; const MANIFEST: &[u8] = br#"{ @@ -1071,6 +1138,26 @@ mod tests { } } + fn spawn_connect(core: *mut CoreHandle) -> thread::JoinHandle { + let address = core as usize; + thread::spawn(move || { + let request = ConnectRequest { + struct_size: u32::try_from(size_of::()).expect("request size"), + abi_version: 1, + manifest_json: BytesView { + data: MANIFEST.as_ptr(), + length: MANIFEST.len(), + }, + tunnel_credential_json: BytesView { + data: CREDENTIAL.as_ptr(), + length: CREDENTIAL.len(), + }, + }; + // Test invariant: the request and static byte views remain valid for the call. + unsafe { super::verse_core_connect_v1(address as *mut CoreHandle, &raw const request) } + }) + } + fn inner(core: *mut CoreHandle) -> Arc { lock(handles()) .get(&(core as usize)) @@ -1204,23 +1291,7 @@ mod tests { resume: resume_rx, }); - let address = core as usize; - let connect = thread::spawn(move || { - let request = ConnectRequest { - struct_size: u32::try_from(size_of::()).expect("request size"), - abi_version: 1, - manifest_json: BytesView { - data: MANIFEST.as_ptr(), - length: MANIFEST.len(), - }, - tunnel_credential_json: BytesView { - data: CREDENTIAL.as_ptr(), - length: CREDENTIAL.len(), - }, - }; - // Test invariant: the request and static byte views remain valid for the call. - unsafe { super::verse_core_connect_v1(address as *mut CoreHandle, &raw const request) } - }); + let connect = spawn_connect(core); reached_rx.recv().expect("final connect reached"); assert_eq!(unsafe { super::verse_core_cancel_v1(core) }, OK); @@ -1238,4 +1309,75 @@ mod tests { ); assert_eq!(connect_status, CANCELLED); } + + #[test] + fn cancellation_during_initial_connect_is_committed_and_emitted() { + let context = Box::new(CallbackContext::default()); + let context_pointer = ptr::from_ref(&*context).cast_mut().cast::(); + let core = create_core_with_callback(context_pointer, Some(record_state)); + let inner = inner(core); + let (reached_tx, reached_rx) = mpsc::channel(); + let (resume_tx, resume_rx) = mpsc::channel(); + *lock(INITIAL_CONNECT_HOOK.get_or_init(|| Mutex::new(None))) = + Some(ConnectTransitionHook { + core: core as usize, + reached: reached_tx, + resume: resume_rx, + }); + let connect = spawn_connect(core); + reached_rx.recv().expect("initial connect reached"); + + assert_eq!(unsafe { super::verse_core_cancel_v1(core) }, OK); + resume_tx.send(()).expect("resume initial connect"); + let connect_status = connect.join().expect("connect caller"); + let cancellation_committed = lock(&inner.state).lifecycle == Lifecycle::Cancelled; + let cancellation_emitted = inner.cancellation_emitted.load(Ordering::Acquire); + if cancellation_emitted { + wait_for_state(&context, STATE_CANCELLED); + } else { + assert_eq!(unsafe { super::verse_core_request_idr_v1(core) }, CANCELLED); + wait_for_state(&context, STATE_CANCELLED); + } + assert_eq!(unsafe { super::verse_core_destroy_v1(core, 2_000) }, OK); + + assert_eq!(connect_status, CANCELLED); + assert!( + cancellation_committed, + "connect returned without committing cancellation" + ); + assert!( + cancellation_emitted, + "connect returned without emitting CANCELLED" + ); + assert_eq!( + context.states.lock().expect("states lock").as_slice(), + [STATE_CONNECTING, STATE_CANCELLED] + ); + } + + #[test] + fn cancellation_wins_the_final_check_to_publish_window() { + let context = Box::new(CallbackContext::default()); + let context_pointer = ptr::from_ref(&*context).cast_mut().cast::(); + let core = create_core_with_callback(context_pointer, Some(record_state)); + let (reached_tx, reached_rx) = mpsc::channel(); + let (resume_tx, resume_rx) = mpsc::channel(); + *lock(FINAL_PUBLISH_HOOK.get_or_init(|| Mutex::new(None))) = Some(ConnectTransitionHook { + core: core as usize, + reached: reached_tx, + resume: resume_rx, + }); + let connect = spawn_connect(core); + reached_rx.recv().expect("final publish reached"); + + assert_eq!(unsafe { super::verse_core_cancel_v1(core) }, OK); + resume_tx.send(()).expect("resume final publish"); + assert_eq!(connect.join().expect("connect caller"), CANCELLED); + wait_for_state(&context, STATE_CANCELLED); + assert_eq!(unsafe { super::verse_core_destroy_v1(core, 2_000) }, OK); + assert_eq!( + context.states.lock().expect("states lock").as_slice(), + [STATE_CONNECTING, STATE_CANCELLED] + ); + } }