diff --git a/tokio/src/task/harness.rs b/tokio/src/task/harness.rs index 0aa9faeeb..6e4555077 100644 --- a/tokio/src/task/harness.rs +++ b/tokio/src/task/harness.rs @@ -300,10 +300,6 @@ where self.drop_waker(); } - pub(super) fn wake_by_local_ref(&self) { - self.wake_by_ref(); - } - pub(super) fn wake_by_ref(&self) { if self.header().state.transition_to_notified() { unsafe { diff --git a/tokio/src/task/tests/task.rs b/tokio/src/task/tests/task.rs index 8f5fec1bf..0cff42950 100644 --- a/tokio/src/task/tests/task.rs +++ b/tokio/src/task/tests/task.rs @@ -641,3 +641,21 @@ fn shutdown_from_task_after_notified() { assert_ready_err!(handle.poll()); } + +#[test] +fn waker_ref_will_wake_clone() { + use std::task::Poll::Ready; + + let (task, handle) = task::joinable(poll_fn(|cx| { + let waker = cx.waker().clone(); + assert!(cx.waker().will_wake(&waker)); + Ready(()) + })); + let mut handle = spawn(handle); + + let mock = mock().bind(&task).release_local(); + let mock = &mut || Some(From::from(&mock)); + + assert_none!(task.run(mock)); + assert_ready_ok!(handle.poll()); +} diff --git a/tokio/src/task/waker.rs b/tokio/src/task/waker.rs index e0e1f36ce..9892f1be8 100644 --- a/tokio/src/task/waker.rs +++ b/tokio/src/task/waker.rs @@ -3,11 +3,12 @@ use crate::task::{Header, Schedule}; use std::future::Future; use std::marker::PhantomData; +use std::mem::ManuallyDrop; use std::ops; use std::task::{RawWaker, RawWakerVTable, Waker}; pub(super) struct WakerRef<'a, S: 'static> { - waker: Waker, + waker: ManuallyDrop, _p: PhantomData<(&'a Header, S)>, } @@ -18,16 +19,15 @@ where T: Future, S: Schedule, { - let ptr = meta as *const _ as *const (); - - let vtable = &RawWakerVTable::new( - clone_waker::, - wake_unreachable, - wake_by_local_ref::, - noop, - ); - - let waker = unsafe { Waker::from_raw(RawWaker::new(ptr, vtable)) }; + // `Waker::will_wake` uses the VTABLE pointer as part of the check. This + // means that `will_wake` will always return false when using the current + // task's waker. (discussion at rust-lang/rust#66281). + // + // To fix this, we use a single vtable. Since we pass in a reference at this + // point and not an *owned* waker, we must ensure that `drop` is never + // called on this waker instance. This is done by wrapping it with + // `ManuallyDrop` and then never calling drop. + let waker = unsafe { ManuallyDrop::new(Waker::from_raw(raw_waker::(meta))) }; WakerRef { waker, @@ -50,15 +50,7 @@ where { let meta = ptr as *const Header; (*meta).state.ref_inc(); - - let vtable = &RawWakerVTable::new( - clone_waker::, - wake_by_val::, - wake_by_ref::, - drop_waker::, - ); - - RawWaker::new(ptr, vtable) + raw_waker::(meta) } unsafe fn drop_waker(ptr: *const ()) @@ -70,11 +62,6 @@ where harness.drop_waker(); } -// `wake()` cannot be called on the ref variaant. -unsafe fn wake_unreachable(_data: *const ()) { - unreachable!(); -} - unsafe fn wake_by_val(ptr: *const ()) where T: Future, @@ -84,16 +71,6 @@ where harness.wake_by_val(); } -// This function can only be called when on the runtime. -unsafe fn wake_by_local_ref(ptr: *const ()) -where - T: Future, - S: Schedule, -{ - let harness = Harness::::from_raw(ptr as *mut _); - harness.wake_by_local_ref(); -} - // Wake without consuming the waker unsafe fn wake_by_ref(ptr: *const ()) where @@ -104,4 +81,17 @@ where harness.wake_by_ref(); } -unsafe fn noop(_ptr: *const ()) {} +fn raw_waker(meta: *const Header) -> RawWaker +where + T: Future, + S: Schedule, +{ + let ptr = meta as *const (); + let vtable = &RawWakerVTable::new( + clone_waker::, + wake_by_val::, + wake_by_ref::, + drop_waker::, + ); + RawWaker::new(ptr, vtable) +}