diff --git a/src/async_event.rs b/src/async_event.rs index d46870b..1dfb454 100644 --- a/src/async_event.rs +++ b/src/async_event.rs @@ -30,18 +30,103 @@ impl AsyncEvent { #[cfg_attr(feature = "profile", hotpath::measure(impl_type = "AsyncEvent"))] pub fn notify_all(&self) { - let mut waiters = self.waiters.lock().expect("poisoned async event waiters"); - let drained = waiters.drain(..); - for waiter in drained { + let mut pending = { + let mut waiters = self.waiters.lock().expect("poisoned async event waiters"); + if waiters.is_empty() { + return; + } + std::mem::take(&mut *waiters) + }; + // Sending can invoke an arbitrary waker synchronously. Notify this + // snapshot outside the lock; reentrant listeners belong to the next + // one. + for waiter in pending.drain(..) { let _ = waiter.send(()); } + // Reuse the allocation when callbacks/concurrent threads have not + // registered more waiters. Never overwrite their live registrations. + let mut waiters = self.waiters.lock().expect("poisoned async event waiters"); + if waiters.is_empty() && pending.capacity() > waiters.capacity() { + *waiters = pending; + } } } #[cfg(test)] mod tests { + use std::{ + future::Future, + pin::Pin, + sync::{ + Arc, Mutex, + atomic::{AtomicBool, Ordering}, + }, + task::{Context, Wake, Waker}, + }; + use super::AsyncEvent; + #[test] + fn notification_callbacks_can_register_the_next_listener() { + struct RegisterListener { + event: Arc, + unlocked: AtomicBool, + next: Mutex>>, + } + impl Wake for RegisterListener { + fn wake(self: Arc) { + self.wake_by_ref(); + } + + fn wake_by_ref(self: &Arc) { + // Detect the old deadlock without making the test hang. + let unlocked = self.event.waiters.try_lock().is_ok(); + self.unlocked.store(unlocked, Ordering::Relaxed); + if unlocked { + *self.next.lock().unwrap() = Some(self.event.listen()); + } + } + } + + let event = Arc::new(AsyncEvent::new()); + let callback = Arc::new(RegisterListener { + event: Arc::clone(&event), + unlocked: AtomicBool::new(false), + next: Mutex::new(None), + }); + let waker = Waker::from(Arc::clone(&callback)); + let mut first = event.listen(); + assert!( + Pin::new(&mut first) + .poll(&mut Context::from_waker(&waker)) + .is_pending() + ); + + event.notify_all(); + assert!(callback.unlocked.load(Ordering::Relaxed)); + assert_eq!(first.try_recv().unwrap(), Some(())); + let mut next = callback.next.lock().unwrap().take().unwrap(); + assert_eq!(next.try_recv().unwrap(), None); + event.notify_all(); + assert_eq!(next.try_recv().unwrap(), Some(())); + } + + #[test] + fn notification_reuses_waiter_capacity_between_bursts() { + let event = AsyncEvent::new(); + let mut listeners: Vec<_> = (0..16).map(|_| event.listen()).collect(); + let capacity = event.waiters.lock().unwrap().capacity(); + event.notify_all(); + assert!( + listeners + .iter_mut() + .all(|listener| listener.try_recv().unwrap() == Some(())) + ); + assert_eq!(event.waiters.lock().unwrap().capacity(), capacity); + event.notify_all(); + assert_eq!(event.waiters.lock().unwrap().capacity(), capacity); + } + #[test] fn notification_wakes_all_current_listeners() { let event = AsyncEvent::new(); diff --git a/src/transport/stream/tls_session.rs b/src/transport/stream/tls_session.rs index 056c02d..ed0aebf 100644 --- a/src/transport/stream/tls_session.rs +++ b/src/transport/stream/tls_session.rs @@ -178,7 +178,14 @@ pub(super) fn complete_tls_handshake( timeout: Duration, server: Option<&Weak>, ) -> io::Result<()> { - let deadline = std::time::Instant::now() + timeout; + let deadline = std::time::Instant::now() + .checked_add(timeout) + .ok_or_else(|| { + io::Error::new( + io::ErrorKind::InvalidInput, + "TLS handshake timeout is too large", + ) + })?; loop { if tls_server_closed(server) { return Err(io::Error::new( @@ -457,6 +464,10 @@ mod tests { .expect_err("zero timeout should fail before handshake I/O"); assert_eq!(err.kind(), io::ErrorKind::TimedOut); + let err = complete_tls_handshake(&tls_state, Duration::MAX, None) + .expect_err("unrepresentable deadline should fail before handshake I/O"); + assert_eq!(err.kind(), io::ErrorKind::InvalidInput); + crate::initialize_python_for_tests(); Python::attach(|py| { let (server, loop_core) = build_test_server(py); diff --git a/src/transport/tls/mod.rs b/src/transport/tls/mod.rs index af81677..92a8f44 100644 --- a/src/transport/tls/mod.rs +++ b/src/transport/tls/mod.rs @@ -137,26 +137,44 @@ pub fn server_tls_settings( }) } +#[derive(Debug, PartialEq, Eq)] +enum TimeoutValueError { + NotPositiveFinite, + TooLarge, +} + +/// Convert scalar input without Python interaction or panicking on overflow. #[cfg_attr(feature = "profile", hotpath::measure)] -fn handshake_timeout(value: Option) -> PyResult { - let secs = value.unwrap_or(DEFAULT_HANDSHAKE_TIMEOUT_SECS); - if !secs.is_finite() || secs <= 0.0 { - return Err(PyValueError::new_err( - "ssl_handshake_timeout must be a positive finite number", - )); +fn timeout_duration(value: f64) -> Result { + if !value.is_finite() || value <= 0.0 { + return Err(TimeoutValueError::NotPositiveFinite); } - Ok(Duration::from_secs_f64(secs)) + Duration::try_from_secs_f64(value).map_err(|_| TimeoutValueError::TooLarge) +} + +#[cfg_attr(feature = "profile", hotpath::measure)] +fn python_timeout(value: Option, default: f64, parameter: &str) -> PyResult { + timeout_duration(value.unwrap_or(default)).map_err(|error| { + let detail = match error { + TimeoutValueError::NotPositiveFinite => "must be a positive finite number", + TimeoutValueError::TooLarge => "is too large", + }; + PyValueError::new_err(format!("{parameter} {detail}")) + }) +} + +#[cfg_attr(feature = "profile", hotpath::measure)] +fn handshake_timeout(value: Option) -> PyResult { + python_timeout( + value, + DEFAULT_HANDSHAKE_TIMEOUT_SECS, + "ssl_handshake_timeout", + ) } #[cfg_attr(feature = "profile", hotpath::measure)] fn shutdown_timeout(value: Option) -> PyResult { - let secs = value.unwrap_or(DEFAULT_SHUTDOWN_TIMEOUT_SECS); - if !secs.is_finite() || secs <= 0.0 { - return Err(PyValueError::new_err( - "ssl_shutdown_timeout must be a positive finite number", - )); - } - Ok(Duration::from_secs_f64(secs)) + python_timeout(value, DEFAULT_SHUTDOWN_TIMEOUT_SECS, "ssl_shutdown_timeout") } #[cfg_attr(feature = "profile", hotpath::measure)] @@ -372,3 +390,60 @@ fn build_server_config(py: Python<'_>, ssl_context: &Py) -> PyResult &'static [&'static SupportedProtocolVersion] { rustls::DEFAULT_VERSIONS } + +#[cfg(test)] +mod timeout_tests { + use super::*; + + #[test] + fn timeout_values_preserve_defaults_rounding_and_invalid_classification() { + assert_eq!(handshake_timeout(None).unwrap(), Duration::from_secs(60)); + assert_eq!(shutdown_timeout(None).unwrap(), Duration::from_secs(30)); + for value in [f64::MIN_POSITIVE, 0.000_000_000_1, 0.25, 1.5, 1e10] { + assert_eq!( + timeout_duration(value).unwrap(), + Duration::from_secs_f64(value) + ); + } + for value in [0.0, -0.0, -1.0, f64::NAN, f64::INFINITY, f64::NEG_INFINITY] { + assert_eq!( + timeout_duration(value), + Err(TimeoutValueError::NotPositiveFinite) + ); + } + assert_eq!(timeout_duration(f64::MAX), Err(TimeoutValueError::TooLarge)); + } + + #[test] + fn invalid_timeout_messages_keep_the_parameter_name() { + crate::initialize_python_for_tests(); + Python::attach(|py| { + type ParseTimeout = fn(Option) -> PyResult; + for (parse, name) in [ + (handshake_timeout as ParseTimeout, "ssl_handshake_timeout"), + (shutdown_timeout, "ssl_shutdown_timeout"), + ] { + let error = parse(Some(-1.0)).unwrap_err(); + assert!(error.is_instance_of::(py)); + assert_eq!( + error.value(py).to_string(), + format!("{name} must be a positive finite number") + ); + } + }); + } + + #[test] + fn oversized_timeouts_report_python_errors_instead_of_panicking() { + crate::initialize_python_for_tests(); + Python::attach(|py| { + for parse in [handshake_timeout, shutdown_timeout] { + let outcome = std::panic::catch_unwind(|| parse(Some(f64::MAX))); + assert!(outcome.is_ok(), "timeout conversion must not panic"); + let error = outcome.unwrap().unwrap_err(); + assert!(error.is_instance_of::(py)); + assert!(error.to_string().contains("too large")); + } + }); + } +} diff --git a/src/vibeio/time/interval.rs b/src/vibeio/time/interval.rs index 3355ebf..e4bd3ec 100644 --- a/src/vibeio/time/interval.rs +++ b/src/vibeio/time/interval.rs @@ -88,80 +88,73 @@ impl Interval { hotpath::measure(impl_type = "Interval", future = true) )] async fn tick_at(&mut self, now: Instant) -> u64 { - // Determine base next (the previous next_deadline or now+period) - let base_next = self - .next_deadline - .unwrap_or_else(|| super::deadline_after(now, self.period)); - - match self.missed_tick_behavior { - MissedTickBehavior::Skip => { - // Advance target forward until it's in the future. - let mut target = base_next; - if target <= now { - if self.period.as_nanos() == 0 { - target = now; - } else { - // Advance directly to the first cadence boundary after - // `now`; iterating once per missed period can otherwise - // make an old interval stall the executor. - let elapsed = now.duration_since(target); - target = super::deadline_after( - now, - self.period - duration_remainder(elapsed, self.period), - ); - } - } - + let (ticks, next_deadline) = match plan_tick( + self.period, + self.next_deadline, + self.missed_tick_behavior, + now, + ) { + TickPlan::WaitUntil(target) => { Sleep::sleep_until_with_zero_behavior(target, super::sleep::ZeroBehavior::Yield) .await; - - // Schedule next deadline for subsequent tick - self.next_deadline = Some(super::deadline_after(target, self.period)); - 1 + (1, super::deadline_after(target, self.period)) } - MissedTickBehavior::CatchUp => { - if base_next > now { - // Not missed yet: sleep until base_next and return 1 - Sleep::sleep_until_with_zero_behavior( - base_next, - super::sleep::ZeroBehavior::Yield, - ) + TickPlan::CatchUp { + ticks, + next_deadline, + } => (ticks, next_deadline), + TickPlan::YieldAndReset => { + Sleep::new_with_zero_behavior(Duration::ZERO, super::sleep::ZeroBehavior::Yield) .await; - self.next_deadline = Some(super::deadline_after(base_next, self.period)); - 1 - } else { - // We missed one or more ticks. Compute how many. - if self.period.as_nanos() == 0 { - // A zero-period catch-up loop must still let other - // tasks run, just as zero-period Skip mode does. - Sleep::new_with_zero_behavior( - Duration::ZERO, - super::sleep::ZeroBehavior::Yield, - ) - .await; - self.next_deadline = Some(Instant::now()); - return 1; - } - - let elapsed = now.duration_since(base_next); - let missed = (elapsed.as_nanos() / self.period.as_nanos()) - .saturating_add(1) - .min(u64::MAX as u128) as u64; - - // The next deadline is the first cadence boundary after - // `now`. `missed` already accounts for the tick at - // `base_next`, so adding another period would skip a tick. - let new_next = super::deadline_after( - now, - self.period - duration_remainder(elapsed, self.period), - ); - self.next_deadline = Some(new_next); - - // Return the number of missed ticks so caller can catch up. - missed - } + (1, Instant::now()) } - } + }; + // Commit only after the wait completes. Dropping a pending tick leaves + // the previous schedule intact, including zero-period catch-up ticks. + self.next_deadline = Some(next_deadline); + ticks + } +} + +/// Scheduling decisions contain no timer registration, clock read or mutation. +#[derive(Debug, PartialEq, Eq)] +enum TickPlan { + WaitUntil(Instant), + CatchUp { ticks: u64, next_deadline: Instant }, + YieldAndReset, +} + +#[cfg_attr(feature = "profile", hotpath::measure)] +#[inline] +fn plan_tick( + period: Duration, + next_deadline: Option, + behavior: MissedTickBehavior, + now: Instant, +) -> TickPlan { + let base_next = next_deadline.unwrap_or_else(|| super::deadline_after(now, period)); + if base_next > now { + return TickPlan::WaitUntil(base_next); + } + if period.is_zero() { + return match behavior { + MissedTickBehavior::Skip => TickPlan::WaitUntil(now), + MissedTickBehavior::CatchUp => TickPlan::YieldAndReset, + }; + } + + let elapsed = now.duration_since(base_next); + // Jump to the first cadence boundary after now, independent of the number + // of missed periods. The remainder preserves sub-millisecond precision. + let next = super::deadline_after(now, period - duration_remainder(elapsed, period)); + match behavior { + MissedTickBehavior::Skip => TickPlan::WaitUntil(next), + MissedTickBehavior::CatchUp => TickPlan::CatchUp { + ticks: (elapsed.as_nanos() / period.as_nanos()) + .saturating_add(1) + .min(u64::MAX as u128) as u64, + next_deadline: next, + }, } } @@ -169,6 +162,89 @@ impl Interval { mod tests { use super::*; + #[test] + fn tick_plan_preserves_cadence_at_and_between_boundaries() { + let base = Instant::now(); + for period_ns in [1, 3, 100_000_000] { + let period = Duration::from_nanos(period_ns); + for elapsed_ns in [0, 1, period_ns - 1, period_ns, 17 * period_ns + 1] { + let now = base + Duration::from_nanos(elapsed_ns); + // An intentionally simple bounded reference: advance one tick + // at a time rather than reuse the production remainder formula. + let mut next = base; + let mut ticks = 0; + while next <= now { + next += period; + ticks += 1; + } + assert_eq!( + plan_tick(period, Some(base), MissedTickBehavior::Skip, now), + TickPlan::WaitUntil(next) + ); + assert_eq!( + plan_tick(period, Some(base), MissedTickBehavior::CatchUp, now), + TickPlan::CatchUp { + ticks, + next_deadline: next + } + ); + } + } + } + + #[test] + fn tick_plan_distinguishes_initial_future_and_zero_period_ticks() { + let now = Instant::now(); + let period = Duration::from_secs(1); + let future = now + period; + for behavior in [MissedTickBehavior::Skip, MissedTickBehavior::CatchUp] { + assert_eq!( + plan_tick(period, None, behavior, now), + TickPlan::WaitUntil(future) + ); + assert_eq!( + plan_tick(period, Some(future), behavior, now), + TickPlan::WaitUntil(future) + ); + // An explicit future first tick must still wait with a zero period. + assert_eq!( + plan_tick(Duration::ZERO, Some(future), behavior, now), + TickPlan::WaitUntil(future) + ); + } + assert_eq!( + plan_tick(Duration::ZERO, None, MissedTickBehavior::Skip, now), + TickPlan::WaitUntil(now) + ); + assert_eq!( + plan_tick(Duration::ZERO, None, MissedTickBehavior::CatchUp, now), + TickPlan::YieldAndReset + ); + } + + #[test] + fn cancelling_future_tick_preserves_explicit_and_initial_schedules() { + Runtime::new(AnyDriver::new_mock()).block_on(async { + let now = Instant::now(); + for behavior in [MissedTickBehavior::Skip, MissedTickBehavior::CatchUp] { + for original in [None, Some(now + Duration::from_secs(60))] { + let mut interval = Interval::new(Duration::from_secs(60)); + interval.set_missed_tick_behavior(behavior); + interval.next_deadline = original; + { + let mut tick = std::pin::pin!(interval.tick_at(now)); + assert!( + tick.as_mut() + .poll(&mut Context::from_waker(Waker::noop())) + .is_pending() + ); + } + assert_eq!(interval.next_deadline, original); + } + } + }); + } + #[test] fn overdue_large_period_saturates_its_next_deadline() { Runtime::new(AnyDriver::new_mock()).block_on(async { diff --git a/tests/test_tls.py b/tests/test_tls.py index f6b5fc1..6089a98 100644 --- a/tests/test_tls.py +++ b/tests/test_tls.py @@ -576,9 +576,11 @@ def test_stream_writer_start_tls_round_trip(tmp_path, backend, client_first): if backend != "rsloop" and not hasattr(asyncio.StreamWriter, "start_tls"): pytest.skip("stdlib StreamWriter.start_tls requires Python 3.11+") timeouts = {"ssl_handshake_timeout": 3} - if backend == "rsloop" or "ssl_shutdown_timeout" in inspect.signature( - asyncio.StreamWriter.start_tls - ).parameters: + if ( + backend == "rsloop" + or "ssl_shutdown_timeout" + in inspect.signature(asyncio.StreamWriter.start_tls).parameters + ): timeouts["ssl_shutdown_timeout"] = 1 factory = ( pytest.importorskip("uvloop").new_event_loop @@ -602,12 +604,7 @@ async def serve(reader, writer): await writer.drain() await server_go.wait() old_transport = writer.transport - assert ( - await writer.start_tls( - server_ctx, **timeouts - ) - is None - ) + assert await writer.start_tls(server_ctx, **timeouts) is None assert writer.transport is not old_transport # Older asyncio versions retain the reader's private transport. if backend == "rsloop": @@ -826,3 +823,18 @@ async def serve(reader, writer): await server.wait_closed() rsloop.run(asyncio.wait_for(main(), 5)) + + +@pytest.mark.parametrize("keyword", ["ssl_handshake_timeout", "ssl_shutdown_timeout"]) +@pytest.mark.parametrize("value", [-1.0, float("nan"), sys.float_info.max]) +def test_tls_timeout_validation_reports_value_error(tmp_path, keyword, value): + server_ctx, _ = make_ssl_contexts(str(tmp_path)) + + async def main(): + loop = asyncio.get_running_loop() + with pytest.raises(ValueError, match=keyword): + await loop.create_server( + asyncio.Protocol, "127.0.0.1", 0, ssl=server_ctx, **{keyword: value} + ) + + rsloop.run(main()) diff --git a/tools/vibeio-check/examples/functional_core.rs b/tools/vibeio-check/examples/functional_core.rs new file mode 100644 index 0000000..d6fd2fa --- /dev/null +++ b/tools/vibeio-check/examples/functional_core.rs @@ -0,0 +1,74 @@ +//! Compare uninstrumented builds: cargo run --release --manifest-path +//! tools/vibeio-check/Cargo.toml --example functional_core -- 100000 7 +use std::{ + future::Future, + hint::black_box, + pin::Pin, + task::{Context, Waker}, + time::{Duration, Instant}, +}; + +use rsloop_vibeio_check::vibeio::{ + DriverKind, RuntimeBuilder, + time::{Interval, MissedTickBehavior}, +}; + +// Exercise the actual private notification implementation, not a copy. +#[path = "../../../src/async_event.rs"] +mod async_event; + +fn notifications(rounds: usize, listeners: usize) -> f64 { + let event = async_event::AsyncEvent::new(); + let mut pending = Vec::with_capacity(listeners); + let mut cx = Context::from_waker(Waker::noop()); + let started = Instant::now(); + for _ in 0..rounds { + for _ in 0..listeners { + let mut listener = event.listen(); + assert!(Pin::new(&mut listener).poll(&mut cx).is_pending()); + pending.push(listener); + } + event.notify_all(); + for mut listener in pending.drain(..) { + assert_eq!(listener.try_recv().unwrap(), Some(())); + } + } + started.elapsed().as_nanos() as f64 / rounds as f64 +} + +fn catch_up_ticks(rounds: usize) -> f64 { + RuntimeBuilder::new() + .driver(DriverKind::Mock) + .enable_timer(true) + .build() + .unwrap() + .block_on(async move { + let mut interval = Interval::new(Duration::from_nanos(1)); + interval.set_missed_tick_behavior(MissedTickBehavior::CatchUp); + let original = Instant::now() - Duration::from_secs(1); + let started = Instant::now(); + for _ in 0..rounds { + interval.next_deadline = Some(original); + assert!(black_box(interval.tick().await) >= 1_000_000_000); + } + started.elapsed().as_nanos() as f64 / rounds as f64 + }) +} + +fn main() { + let mut args = std::env::args().skip(1); + let rounds: usize = args.next().map_or(100_000, |n| n.parse().unwrap()); + let samples: usize = args.next().map_or(7, |n| n.parse().unwrap()); + assert!(rounds > 0 && samples > 0); + for listeners in [0, 1, 4, 16] { + notifications(1000, listeners); + for sample in 0..samples { + let elapsed = notifications(rounds, listeners); + println!("event-{listeners},{sample},{elapsed:.2}"); + } + } + catch_up_ticks(1000); + for sample in 0..samples { + println!("catch-up,{sample},{:.2}", catch_up_ticks(rounds)); + } +}