use smallvec::SmallVec;
const SLOT_BITS: u32 = 16;
const SLOT_COUNT: usize = 1 << SLOT_BITS;
const SLOT_MASK: u64 = (SLOT_COUNT as u64) - 1;
const MAX_TIMERS: usize = 65536;
const FREE: u32 = u32::MAX;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TimerType {
TcpRetransmit,
TcpConnectionTimeout,
TcpTimeWait,
TcpFinWait2,
UdpSessionTimeout,
QuicPto,
QuicCloseTimeout,
QuicHandshakeTimeout,
}
pub const TIMER_TYPE_COUNT: usize = 8;
impl TimerType {
#[inline]
pub const fn as_index(self) -> usize {
match self {
TimerType::TcpRetransmit => 0,
TimerType::TcpConnectionTimeout => 1,
TimerType::TcpTimeWait => 2,
TimerType::TcpFinWait2 => 3,
TimerType::UdpSessionTimeout => 4,
TimerType::QuicPto => 5,
TimerType::QuicCloseTimeout => 6,
TimerType::QuicHandshakeTimeout => 7,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TimerAction {
pub timer_type: TimerType,
pub target_idx: usize,
}
#[derive(Clone)]
struct TimerNode {
expires_at: u64,
inserted_at: u64,
timer_type: TimerType,
target_idx: usize,
next: u32,
prev: u32,
active: bool,
}
impl core::fmt::Debug for TimerNode {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("TimerNode")
.field("expires_at", &self.expires_at)
.field("timer_type", &self.timer_type)
.field("target_idx", &self.target_idx)
.finish()
}
}
impl TimerNode {
const fn empty() -> Self {
Self {
expires_at: 0,
inserted_at: 0,
timer_type: TimerType::UdpSessionTimeout,
target_idx: 0,
next: FREE,
prev: FREE,
active: false,
}
}
}
pub struct TimerWheel {
slots: Vec<u32>,
nodes: Vec<TimerNode>,
free_head: u32,
free_count: u32,
current_tick: u64,
tick_ms: u64,
}
impl core::fmt::Debug for TimerWheel {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("TimerWheel")
.field("current_tick", &self.current_tick)
.field("tick_ms", &self.tick_ms)
.field("active", &(MAX_TIMERS as u32 - self.free_count))
.finish()
}
}
impl TimerWheel {
pub fn new(tick_ms: u64) -> Self {
let mut nodes = Vec::with_capacity(MAX_TIMERS);
for i in 0..MAX_TIMERS {
let mut node = TimerNode::empty();
node.next = if i + 1 < MAX_TIMERS {
(i + 1) as u32
} else {
FREE
};
nodes.push(node);
}
Self {
slots: vec![FREE; SLOT_COUNT],
nodes,
free_head: 0,
free_count: MAX_TIMERS as u32,
current_tick: 0,
tick_ms,
}
}
#[inline]
pub fn current_tick(&self) -> u64 {
self.current_tick
}
#[inline]
pub fn tick_ms(&self) -> u64 {
self.tick_ms
}
#[inline]
pub fn active_count(&self) -> u32 {
(MAX_TIMERS as u32) - self.free_count
}
pub fn add_timer(
&mut self,
duration_ms: u64,
timer_type: TimerType,
target_idx: usize,
) -> Option<u32> {
if self.free_count == 0 {
return None;
}
let node_idx = self.free_head;
self.free_head = self.nodes[node_idx as usize].next;
self.free_count -= 1;
let duration = duration_ms.max(1);
let Some(expires_at) = self.current_tick.checked_add(duration) else {
self.free_node(node_idx);
return None;
};
let expires_at = if duration > SLOT_COUNT as u64 {
tracing::warn!(
duration_ms,
max_ms = SLOT_COUNT as u64 * self.tick_ms,
"定时器延迟超过时间轮容量,已钳制到最大延迟"
);
self.current_tick + (SLOT_COUNT as u64 - 1)
} else {
expires_at
};
let slot = (expires_at & SLOT_MASK) as usize;
self.nodes[node_idx as usize].expires_at = expires_at;
self.nodes[node_idx as usize].inserted_at = self.current_tick;
self.nodes[node_idx as usize].timer_type = timer_type;
self.nodes[node_idx as usize].target_idx = target_idx;
self.nodes[node_idx as usize].active = true;
let old_head = self.slots[slot];
self.nodes[node_idx as usize].next = old_head;
self.nodes[node_idx as usize].prev = FREE;
if old_head != FREE {
self.nodes[old_head as usize].prev = node_idx;
}
self.slots[slot] = node_idx;
Some(node_idx)
}
pub fn remove_timer(&mut self, target_idx: usize, timer_type: TimerType) -> bool {
for slot in 0..SLOT_COUNT {
let mut ni = self.slots[slot];
while ni != FREE {
let next = self.nodes[ni as usize].next;
if self.nodes[ni as usize].target_idx == target_idx
&& self.nodes[ni as usize].timer_type == timer_type
{
self.unlink_node(ni, slot);
self.free_node(ni);
return true;
}
ni = next;
}
}
false
}
pub fn remove_by_handle(&mut self, handle: u32) -> Option<TimerAction> {
if (handle as usize) >= MAX_TIMERS || !self.nodes[handle as usize].active {
return None;
}
let expires_at = self.nodes[handle as usize].expires_at;
let slot = (expires_at & SLOT_MASK) as usize;
let mut ni = self.slots[slot];
while ni != FREE {
let next = self.nodes[ni as usize].next;
if ni == handle {
let action = TimerAction {
timer_type: self.nodes[ni as usize].timer_type,
target_idx: self.nodes[ni as usize].target_idx,
};
self.unlink_node(ni, slot);
self.free_node(ni);
return Some(action);
}
ni = next;
}
None
}
pub fn advance(&mut self, elapsed_ms: u64) -> SmallVec<[TimerAction; 16]> {
let mut expired = SmallVec::new();
let mut remaining = elapsed_ms;
while remaining >= self.tick_ms {
self.current_tick += 1;
remaining -= self.tick_ms;
self.process_current_slot(&mut expired);
}
expired
}
fn process_current_slot(&mut self, expired: &mut SmallVec<[TimerAction; 16]>) {
let slot = (self.current_tick & SLOT_MASK) as usize;
let mut ni = self.slots[slot];
self.slots[slot] = FREE;
while ni != FREE {
let next = self.nodes[ni as usize].next;
if self.nodes[ni as usize].expires_at <= self.current_tick {
let action = TimerAction {
timer_type: self.nodes[ni as usize].timer_type,
target_idx: self.nodes[ni as usize].target_idx,
};
expired.push(action);
self.free_node(ni);
} else {
self.reinsert_node(ni);
}
ni = next;
}
}
fn reinsert_node(&mut self, node_idx: u32) {
let expires_at = self.nodes[node_idx as usize].expires_at;
let new_slot = (expires_at & SLOT_MASK) as usize;
self.nodes[node_idx as usize].next = self.slots[new_slot];
self.nodes[node_idx as usize].prev = FREE;
if self.slots[new_slot] != FREE {
self.nodes[self.slots[new_slot] as usize].prev = node_idx;
}
self.slots[new_slot] = node_idx;
}
fn unlink_node(&mut self, node_idx: u32, slot: usize) {
let prev = self.nodes[node_idx as usize].prev;
let next = self.nodes[node_idx as usize].next;
if prev != FREE {
self.nodes[prev as usize].next = next;
} else {
self.slots[slot] = next;
}
if next != FREE {
self.nodes[next as usize].prev = prev;
}
self.nodes[node_idx as usize].next = FREE;
self.nodes[node_idx as usize].prev = FREE;
}
fn free_node(&mut self, node_idx: u32) {
self.nodes[node_idx as usize].active = false;
self.nodes[node_idx as usize].next = self.free_head;
self.free_head = node_idx;
self.free_count += 1;
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_single_timer_expire() {
let mut wheel = TimerWheel::new(1);
wheel.add_timer(10, TimerType::TcpRetransmit, 0);
assert_eq!(wheel.active_count(), 1);
let expired = wheel.advance(10);
assert_eq!(expired.len(), 1);
assert_eq!(expired[0].timer_type, TimerType::TcpRetransmit);
assert_eq!(expired[0].target_idx, 0);
assert_eq!(wheel.active_count(), 0);
}
#[test]
fn test_multiple_timers_different_delays() {
let mut wheel = TimerWheel::new(1);
wheel.add_timer(10, TimerType::TcpRetransmit, 0);
wheel.add_timer(20, TimerType::UdpSessionTimeout, 1);
wheel.add_timer(5, TimerType::TcpTimeWait, 2);
assert_eq!(wheel.active_count(), 3);
let expired = wheel.advance(5);
assert_eq!(expired.len(), 1);
assert_eq!(expired[0].timer_type, TimerType::TcpTimeWait);
assert_eq!(wheel.active_count(), 2);
let expired = wheel.advance(5);
assert_eq!(expired.len(), 1);
assert_eq!(expired[0].target_idx, 0);
assert_eq!(wheel.active_count(), 1);
let expired = wheel.advance(10);
assert_eq!(expired.len(), 1);
assert_eq!(expired[0].target_idx, 1);
assert_eq!(wheel.active_count(), 0);
}
#[test]
fn test_remove_by_handle() {
let mut wheel = TimerWheel::new(1);
let h = wheel.add_timer(100, TimerType::TcpConnectionTimeout, 42).unwrap();
assert_eq!(wheel.active_count(), 1);
let action = wheel.remove_by_handle(h).unwrap();
assert_eq!(action.target_idx, 42);
assert_eq!(wheel.active_count(), 0);
assert!(wheel.remove_by_handle(h).is_none());
}
#[test]
fn test_remove_by_target() {
let mut wheel = TimerWheel::new(1);
wheel.add_timer(100, TimerType::TcpRetransmit, 5);
wheel.add_timer(200, TimerType::UdpSessionTimeout, 10);
assert_eq!(wheel.active_count(), 2);
let removed = wheel.remove_timer(5, TimerType::TcpRetransmit);
assert!(removed);
assert_eq!(wheel.active_count(), 1);
let expired = wheel.advance(300);
assert_eq!(expired.len(), 1);
assert_eq!(expired[0].target_idx, 10);
}
#[test]
fn test_many_timers_same_slot() {
let mut wheel = TimerWheel::new(1);
for i in 0..10u32 {
wheel.add_timer(50, TimerType::UdpSessionTimeout, i as usize);
}
assert_eq!(wheel.active_count(), 10);
let expired = wheel.advance(50);
assert_eq!(expired.len(), 10);
assert_eq!(wheel.active_count(), 0);
}
#[test]
fn test_wraparound() {
let mut wheel = TimerWheel::new(1);
wheel.add_timer(70000, TimerType::TcpTimeWait, 99);
assert_eq!(wheel.active_count(), 1);
for _ in 0..70 {
wheel.advance(1000);
}
assert_eq!(wheel.active_count(), 0);
}
#[test]
fn test_bulk_operations() {
let mut wheel = TimerWheel::new(1);
for i in 0..500 {
let delay = (i % 100) as u64 + 1;
wheel.add_timer(delay, TimerType::UdpSessionTimeout, i);
}
assert_eq!(wheel.active_count(), 500);
let mut total = 0;
for _ in 0..200 {
total += wheel.advance(1).len();
}
assert_eq!(total, 500);
assert_eq!(wheel.active_count(), 0);
}
#[test]
fn test_multiple_fifo_same_slot() {
let mut wheel = TimerWheel::new(1);
wheel.add_timer(10, TimerType::TcpRetransmit, 0);
wheel.add_timer(10, TimerType::TcpRetransmit, 1);
wheel.add_timer(10, TimerType::TcpRetransmit, 2);
let expired = wheel.advance(10);
assert_eq!(expired.len(), 3);
assert_eq!(expired[0].target_idx, 2);
assert_eq!(expired[1].target_idx, 1);
assert_eq!(expired[2].target_idx, 0);
}
#[test]
fn test_timer_wheel_initial_state() {
let wheel = TimerWheel::new(10);
assert_eq!(wheel.current_tick(), 0);
assert_eq!(wheel.tick_ms(), 10);
assert_eq!(wheel.active_count(), 0);
}
#[test]
fn test_timer_zero_duration() {
let mut wheel = TimerWheel::new(1);
let h = wheel.add_timer(0, TimerType::UdpSessionTimeout, 0).unwrap();
assert_eq!(wheel.active_count(), 1);
let expired = wheel.advance(1);
assert_eq!(expired.len(), 1);
assert_eq!(expired[0].target_idx, 0);
let _ = h;
}
#[test]
fn test_timer_remove_by_handle_invalid() {
let mut wheel = TimerWheel::new(1);
assert!(wheel.remove_by_handle(u32::MAX).is_none());
assert!(wheel.remove_by_handle(0).is_none());
}
#[test]
fn test_timer_remove_by_target_not_found() {
let mut wheel = TimerWheel::new(1);
assert!(!wheel.remove_timer(999, TimerType::TcpRetransmit));
}
#[test]
fn test_timer_wheel_65536_wraparound() {
let mut wheel = TimerWheel::new(1);
wheel.add_timer(65536, TimerType::TcpTimeWait, 42);
assert_eq!(wheel.active_count(), 1);
for _ in 0..65536 {
wheel.advance(1);
}
assert_eq!(wheel.active_count(), 0);
assert_eq!(wheel.current_tick(), 65536);
}
#[test]
fn test_timer_wheel_slot_zero() {
let mut wheel = TimerWheel::new(1);
wheel.add_timer(1, TimerType::UdpSessionTimeout, 100);
assert_eq!(wheel.active_count(), 1);
let expired = wheel.advance(1);
assert_eq!(expired.len(), 1);
assert_eq!(expired[0].target_idx, 100);
}
#[test]
fn test_timer_wheel_max_slot() {
let mut wheel = TimerWheel::new(1);
wheel.add_timer(65535, TimerType::TcpFinWait2, 200);
assert_eq!(wheel.active_count(), 1);
let expired = wheel.advance(65535);
assert_eq!(expired.len(), 1);
assert_eq!(expired[0].target_idx, 200);
}
#[test]
fn test_timer_multiple_advances() {
let mut wheel = TimerWheel::new(1);
wheel.add_timer(100, TimerType::TcpRetransmit, 0);
wheel.add_timer(200, TimerType::UdpSessionTimeout, 1);
let mut total_expired = 0;
for _ in 0..50 {
total_expired += wheel.advance(5).len();
}
assert_eq!(total_expired, 2);
assert_eq!(wheel.active_count(), 0);
}
#[test]
fn test_timer_action_fields() {
let action = TimerAction {
timer_type: TimerType::TcpRetransmit,
target_idx: 42,
};
assert_eq!(action.timer_type, TimerType::TcpRetransmit);
assert_eq!(action.target_idx, 42);
}
#[test]
fn test_timer_type_all_variants() {
let types = vec![
TimerType::TcpRetransmit,
TimerType::TcpConnectionTimeout,
TimerType::TcpTimeWait,
TimerType::TcpFinWait2,
TimerType::UdpSessionTimeout,
TimerType::QuicPto,
TimerType::QuicCloseTimeout,
TimerType::QuicHandshakeTimeout,
];
for t in &types {
let _ = format!("{:?}", t);
}
}
#[test]
fn test_timer_node_debug() {
let node = TimerNode::empty();
let debug = format!("{:?}", node);
assert!(debug.contains("TimerNode"));
}
#[test]
fn test_timer_wheel_debug() {
let wheel = TimerWheel::new(100);
let debug = format!("{:?}", wheel);
assert!(debug.contains("TimerWheel"));
assert!(debug.contains("current_tick"));
}
#[test]
fn test_bulk_expire_many_timers() {
let mut wheel = TimerWheel::new(1);
for i in 0..100u32 {
wheel.add_timer(50, TimerType::UdpSessionTimeout, i as usize);
}
assert_eq!(wheel.active_count(), 100);
let expired = wheel.advance(50);
assert_eq!(expired.len(), 100);
assert_eq!(wheel.active_count(), 0);
}
#[test]
fn test_add_after_expiry() {
let mut wheel = TimerWheel::new(1);
wheel.add_timer(10, TimerType::TcpRetransmit, 0);
wheel.advance(10);
assert_eq!(wheel.active_count(), 0);
wheel.add_timer(5, TimerType::UdpSessionTimeout, 1);
assert_eq!(wheel.active_count(), 1);
let expired = wheel.advance(5);
assert_eq!(expired.len(), 1);
assert_eq!(expired[0].target_idx, 1);
}
}