kestrel_timer/service.rs
1use crate::config::ServiceConfig;
2use crate::error::TimerError;
3use crate::task::{
4 CallbackWrapper, CompletionReceiver, TaskCompletion, TaskHandle, TaskId, TimerTask,
5 TimerTaskWithCompletionNotifier,
6};
7use crate::wheel::Wheel;
8use crate::{BatchHandle, TimerHandle};
9use futures::future::BoxFuture;
10use futures::stream::{FuturesUnordered, StreamExt};
11use lite_sync::{
12 oneshot::lite::{Receiver, Sender, channel},
13 spsc,
14};
15use parking_lot::Mutex;
16use std::collections::HashSet;
17use std::sync::Arc;
18use std::time::Duration;
19use tokio::sync::mpsc;
20use tokio::sync::watch;
21use tokio::task::JoinHandle;
22
23/// Task notification type for distinguishing between one-shot and periodic tasks
24///
25/// 任务通知类型,用于区分一次性任务和周期性任务
26#[derive(Debug, Clone, Copy, PartialEq, Eq)]
27pub enum TaskNotification {
28 /// One-shot task expired notification
29 ///
30 /// 一次性任务过期通知
31 OneShot(TaskId),
32 /// Periodic task called notification
33 ///
34 /// 周期性任务被调用通知
35 Periodic(TaskId),
36}
37
38impl TaskNotification {
39 /// Get the task ID from the notification
40 ///
41 /// 从通知中获取任务 ID
42 pub fn task_id(&self) -> TaskId {
43 match self {
44 TaskNotification::OneShot(id) => *id,
45 TaskNotification::Periodic(id) => *id,
46 }
47 }
48
49 /// Check if this is a one-shot task notification
50 ///
51 /// 检查是否为一次性任务通知
52 pub fn is_oneshot(&self) -> bool {
53 matches!(self, TaskNotification::OneShot(_))
54 }
55
56 /// Check if this is a periodic task notification
57 ///
58 /// 检查是否为周期性任务通知
59 pub fn is_periodic(&self) -> bool {
60 matches!(self, TaskNotification::Periodic(_))
61 }
62}
63
64/// Service command type
65///
66/// 服务命令类型
67enum ServiceCommand {
68 /// Add batch timer handle, only contains necessary data: task_ids and completion_rxs
69 ///
70 /// 添加批量定时器句柄,仅包含必要数据:task_ids 和 completion_rxs
71 AddBatchHandle {
72 task_ids: Vec<TaskId>,
73 completion_rxs: Vec<CompletionReceiver>,
74 },
75 /// Add single timer handle, only contains necessary data: task_id and completion_rx
76 ///
77 /// 添加单个定时器句柄,仅包含必要数据:task_id 和 completion_rx
78 AddTimerHandle {
79 task_id: TaskId,
80 completion_rx: CompletionReceiver,
81 },
82}
83
84/// TimerService - timer service based on Actor pattern
85/// Manages multiple timer handles, listens to all timeout events, and aggregates notifications to be forwarded to the user.
86/// # Features
87/// - Automatically listens to all added timer handles' timeout events
88/// - Automatically removes one-shot tasks from internal management after timeout
89/// - Continuously monitors periodic tasks and forwards each invocation
90/// - Aggregates notifications (both one-shot and periodic) to be forwarded to the user's unified channel
91/// - Supports dynamic addition of BatchHandle and TimerHandle
92///
93///
94/// # 定时器服务,基于 Actor 模式管理多个定时器句柄,监听所有超时事件,并将通知聚合转发给用户
95/// - 自动监听所有添加的定时器句柄的超时事件
96/// - 自动在一次性任务超时后从内部管理中移除任务
97/// - 持续监听周期性任务并转发每次调用通知
98/// - 将通知(一次性和周期性)聚合转发给用户
99/// - 支持动态添加 BatchHandle 和 TimerHandle
100///
101/// # Examples (示例)
102/// ```no_run
103/// use kestrel_timer::{TimerWheel, TimerService, CallbackWrapper, TaskNotification, config::ServiceConfig};
104/// use std::time::Duration;
105///
106/// #[tokio::main]
107/// async fn main() {
108/// let timer = TimerWheel::with_defaults();
109/// let mut service = timer.create_service(ServiceConfig::default());
110///
111/// // Register one-shot tasks (注册一次性任务)
112/// use kestrel_timer::TimerTask;
113/// let handles = service.allocate_handles(3);
114/// let tasks: Vec<_> = (0..3)
115/// .map(|i| {
116/// let callback = Some(CallbackWrapper::new(move || async move {
117/// println!("One-shot timer {} fired!", i);
118/// }));
119/// TimerTask::new_oneshot(Duration::from_millis(100), callback)
120/// })
121/// .collect();
122/// service.register_batch(handles, tasks).unwrap();
123///
124/// // Register periodic tasks (注册周期性任务)
125/// let handle = service.allocate_handle();
126/// let periodic_task = TimerTask::new_periodic(
127/// Duration::from_millis(100),
128/// Duration::from_millis(50),
129/// Some(CallbackWrapper::new(|| async { println!("Periodic timer fired!"); })),
130/// None
131/// )
132/// .unwrap();
133/// service.register(handle, periodic_task).unwrap();
134///
135/// // Receive notifications (接收通知)
136/// let rx = service.take_receiver().unwrap();
137/// while let Some(notification) = rx.recv().await {
138/// match notification {
139/// TaskNotification::OneShot(task_id) => {
140/// println!("One-shot task {:?} expired", task_id);
141/// }
142/// TaskNotification::Periodic(task_id) => {
143/// println!("Periodic task {:?} called", task_id);
144/// }
145/// }
146/// }
147/// }
148/// ```
149pub struct TimerService {
150 /// Command sender
151 ///
152 /// 命令发送器
153 command_tx: mpsc::Sender<ServiceCommand>,
154 /// Timeout receiver (supports both one-shot and periodic task notifications)
155 ///
156 /// 超时接收器(支持一次性和周期性任务通知)
157 timeout_rx: Option<spsc::Receiver<TaskNotification, 32>>,
158 /// Actor task handle
159 ///
160 /// Actor 任务句柄
161 actor_handle: Option<JoinHandle<()>>,
162 /// Timing wheel reference (for direct scheduling of timers)
163 ///
164 /// 时间轮引用(用于直接调度定时器)
165 wheel: Arc<Mutex<Wheel>>,
166 /// Actor shutdown signal sender
167 ///
168 /// Actor 关闭信号发送器
169 shutdown_tx: Option<Sender<()>>,
170 /// Tasks registered through this service and still eligible for cleanup
171 ///
172 /// 通过此服务注册、仍需由服务清理的任务
173 registered_task_ids: Arc<Mutex<HashSet<TaskId>>>,
174}
175
176impl TimerService {
177 /// Allocate a handle from DeferredMap
178 ///
179 /// # Returns
180 /// A unique handle for later insertion
181 ///
182 /// # 返回值
183 /// 用于后续插入的唯一 handle
184 ///
185 /// # Examples (示例)
186 /// ```no_run
187 /// # use kestrel_timer::{TimerWheel, config::ServiceConfig};
188 /// # #[tokio::main]
189 /// # async fn main() {
190 /// let timer = TimerWheel::with_defaults();
191 /// let mut service = timer.create_service(ServiceConfig::default());
192 ///
193 /// // Allocate handle first
194 /// // 先分配handle
195 /// let handle = service.allocate_handle();
196 /// # }
197 /// ```
198 pub fn allocate_handle(&self) -> TaskHandle {
199 self.wheel.lock().allocate_handle()
200 }
201
202 /// Batch allocate handles from DeferredMap
203 ///
204 /// # Parameters
205 /// - `count`: Number of handles to allocate
206 ///
207 /// # Returns
208 /// Vector of unique handles for later batch insertion
209 ///
210 /// # 参数
211 /// - `count`: 要分配的 handle 数量
212 ///
213 /// # 返回值
214 /// 用于后续批量插入的唯一 handles 向量
215 ///
216 /// # Examples (示例)
217 /// ```no_run
218 /// # use kestrel_timer::{TimerWheel, config::ServiceConfig};
219 /// # #[tokio::main]
220 /// # async fn main() {
221 /// let timer = TimerWheel::with_defaults();
222 /// let service = timer.create_service(ServiceConfig::default());
223 ///
224 /// // Batch allocate handles
225 /// // 批量分配 handles
226 /// let handles = service.allocate_handles(10);
227 /// assert_eq!(handles.len(), 10);
228 /// # }
229 /// ```
230 pub fn allocate_handles(&self, count: usize) -> Vec<TaskHandle> {
231 self.wheel.lock().allocate_handles(count)
232 }
233
234 /// Create new TimerService
235 ///
236 /// # Parameters
237 /// - `wheel`: Timing wheel reference
238 /// - `config`: Service configuration
239 ///
240 /// # Notes
241 /// Typically not called directly, but used to create through `TimerWheel::create_service()`
242 ///
243 /// 创建新的 TimerService
244 ///
245 /// # 参数
246 /// - `wheel`: 时间轮引用
247 /// - `config`: 服务配置
248 ///
249 /// # 注意
250 /// 通常不直接调用,而是通过 `TimerWheel::create_service()` 创建
251 ///
252 pub(crate) fn new(wheel: Arc<Mutex<Wheel>>, config: ServiceConfig) -> Self {
253 let (command_tx, command_rx) = mpsc::channel(config.command_channel_capacity.get());
254 let (timeout_tx, timeout_rx) = spsc::channel(config.timeout_channel_capacity);
255
256 let (shutdown_tx, shutdown_rx) = channel::<()>();
257 let registered_task_ids = Arc::new(Mutex::new(HashSet::new()));
258 let wheel_closed = wheel.lock().owner().subscribe_closed();
259 let actor = ServiceActor::new(
260 command_rx,
261 timeout_tx,
262 shutdown_rx,
263 wheel_closed,
264 Arc::clone(®istered_task_ids),
265 );
266 let actor_handle = tokio::spawn(async move {
267 actor.run().await;
268 });
269
270 Self {
271 command_tx,
272 timeout_rx: Some(timeout_rx),
273 actor_handle: Some(actor_handle),
274 wheel,
275 shutdown_tx: Some(shutdown_tx),
276 registered_task_ids,
277 }
278 }
279
280 /// Get timeout receiver (transfer ownership)
281 ///
282 /// # Returns
283 /// Timeout notification receiver, if already taken, returns None
284 ///
285 /// # Notes
286 /// This method can only be called once, because it transfers ownership of the receiver
287 /// The receiver will receive both one-shot task expired notifications and periodic task called notifications
288 ///
289 /// 获取超时通知接收器(转移所有权)
290 ///
291 /// # 返回值
292 /// 超时通知接收器,如果已经取走,返回 None
293 ///
294 /// # 注意
295 /// 此方法只能调用一次,因为它转移了接收器的所有权
296 /// 接收器将接收一次性任务过期通知和周期性任务被调用通知
297 ///
298 /// # Examples (示例)
299 /// ```no_run
300 /// # use kestrel_timer::{TimerWheel, config::ServiceConfig, TaskNotification};
301 /// # #[tokio::main]
302 /// # async fn main() {
303 /// let timer = TimerWheel::with_defaults();
304 /// let mut service = timer.create_service(ServiceConfig::default());
305 ///
306 /// let rx = service.take_receiver().unwrap();
307 /// while let Some(notification) = rx.recv().await {
308 /// match notification {
309 /// TaskNotification::OneShot(task_id) => {
310 /// println!("One-shot task {:?} expired", task_id);
311 /// }
312 /// TaskNotification::Periodic(task_id) => {
313 /// println!("Periodic task {:?} called", task_id);
314 /// }
315 /// }
316 /// }
317 /// # }
318 /// ```
319 pub fn take_receiver(&mut self) -> Option<spsc::Receiver<TaskNotification, 32>> {
320 self.timeout_rx.take()
321 }
322
323 /// Cancel specified task
324 ///
325 /// # Parameters
326 /// - `task_id`: Task ID to cancel
327 ///
328 /// # Returns
329 /// - `Ok(true)`: Task exists and cancellation is successful
330 /// - `Ok(false)`: Task does not exist or cancellation fails
331 /// - `Err(TimerError::WrongWheel)`: Task ID belongs to another wheel
332 ///
333 /// 取消指定任务
334 ///
335 /// # 参数
336 /// - `task_id`: 任务 ID
337 ///
338 /// # 返回值
339 /// - `Ok(true)`: 任务存在且取消成功
340 /// - `Ok(false)`: 任务不存在或取消失败
341 /// - `Err(TimerError::WrongWheel)`: 任务 ID 属于另一个时间轮
342 ///
343 /// # Examples (示例)
344 /// ```no_run
345 /// # use kestrel_timer::{TimerWheel, TimerService, CallbackWrapper, TimerTask, config::ServiceConfig};
346 /// # use std::time::Duration;
347 /// #
348 /// # #[tokio::main]
349 /// # async fn main() {
350 /// let timer = TimerWheel::with_defaults();
351 /// let service = timer.create_service(ServiceConfig::default());
352 ///
353 /// // Use two-step API to schedule timers
354 /// let handle = service.allocate_handle();
355 /// let task_id = handle.task_id();
356 /// let callback = Some(CallbackWrapper::new(|| async move {
357 /// println!("Timer fired!"); // 定时器触发
358 /// }));
359 /// let task = TimerTask::new_oneshot(Duration::from_secs(10), callback);
360 /// service.register(handle, task).unwrap(); // 注册定时器
361 ///
362 /// // Cancel task
363 /// let cancelled = service.cancel_task(task_id).unwrap();
364 /// println!("Task cancelled: {}", cancelled); // 任务取消
365 /// # }
366 /// ```
367 #[inline]
368 pub fn cancel_task(&self, task_id: TaskId) -> Result<bool, TimerError> {
369 // Direct cancellation, no need to notify Actor
370 // FuturesUnordered will automatically clean up when tasks are cancelled
371 // 直接取消,无需通知 Actor
372 // FuturesUnordered 将在任务取消时自动清理
373 let mut wheel = self.wheel.lock();
374 wheel.cancel(task_id)
375 }
376
377 /// Batch cancel tasks
378 ///
379 /// Use underlying batch cancellation operation to cancel multiple tasks at once, performance is better than calling cancel_task repeatedly.
380 ///
381 /// # Parameters
382 /// - `task_ids`: List of task IDs to cancel
383 ///
384 /// # Returns
385 /// Number of successfully cancelled tasks, or
386 /// `Err(TimerError::WrongWheel)` if any ID belongs to another wheel.
387 ///
388 /// 批量取消任务
389 ///
390 /// # 参数
391 /// - `task_ids`: 任务 ID 列表
392 ///
393 /// # 返回值
394 /// 成功取消的任务数量;如果任一 ID 属于其他时间轮则返回
395 /// `Err(TimerError::WrongWheel)`,且不修改任何任务。
396 ///
397 /// # Examples (示例)
398 /// ```no_run
399 /// # use kestrel_timer::{TimerWheel, TimerService, CallbackWrapper, TimerTask, config::ServiceConfig};
400 /// # use std::time::Duration;
401 /// #
402 /// # #[tokio::main]
403 /// # async fn main() {
404 /// let timer = TimerWheel::with_defaults();
405 /// let service = timer.create_service(ServiceConfig::default());
406 ///
407 /// let handles = service.allocate_handles(10);
408 /// let task_ids: Vec<_> = handles.iter().map(|h| h.task_id()).collect();
409 /// let tasks: Vec<_> = (0..10)
410 /// .map(|i| {
411 /// let callback = Some(CallbackWrapper::new(move || async move {
412 /// println!("Timer {} fired!", i); // 定时器触发
413 /// }));
414 /// TimerTask::new_oneshot(Duration::from_secs(10), callback)
415 /// })
416 /// .collect();
417 /// service.register_batch(handles, tasks).unwrap(); // 注册定时器
418 ///
419 /// // Batch cancel
420 /// let cancelled = service.cancel_batch(&task_ids).unwrap();
421 /// println!("Cancelled {} tasks", cancelled); // 任务取消
422 /// # }
423 /// ```
424 #[inline]
425 pub fn cancel_batch(&self, task_ids: &[TaskId]) -> Result<usize, TimerError> {
426 if task_ids.is_empty() {
427 return Ok(0);
428 }
429
430 // Direct batch cancellation, no need to notify Actor
431 // FuturesUnordered will automatically clean up when tasks are cancelled
432 // 直接批量取消,无需通知 Actor
433 // FuturesUnordered 将在任务取消时自动清理
434 let mut wheel = self.wheel.lock();
435 wheel.cancel_batch(task_ids)
436 }
437
438 /// Postpone task (optionally replace callback)
439 ///
440 /// # Parameters
441 /// - `task_id`: Task ID to postpone
442 /// - `new_delay`: New delay time (recalculated from current time point)
443 /// - `callback`: New callback function (if `None`, keeps the original callback)
444 ///
445 /// # Returns
446 /// - `Ok(true)`: Task exists and is successfully postponed
447 /// - `Ok(false)`: Task does not exist or postponement fails
448 /// - `Err(TimerError::WrongWheel)`: Task ID belongs to another wheel
449 ///
450 /// # Notes
451 /// - Task ID remains unchanged after postponement
452 /// - Original timeout notification remains valid
453 /// - If callback is `Some`, it will replace the original callback
454 /// - If callback is `None`, the original callback is preserved
455 ///
456 /// 推迟任务 (可选替换回调)
457 ///
458 /// # 参数
459 /// - `task_id`: 任务 ID
460 /// - `new_delay`: 新的延迟时间 (从当前时间点重新计算)
461 /// - `callback`: 新的回调函数 (如果为 `None`,则保留原回调)
462 ///
463 /// # 返回值
464 /// - `Ok(true)`: 任务存在且延期成功
465 /// - `Ok(false)`: 任务不存在或延期失败
466 /// - `Err(TimerError::WrongWheel)`: 任务 ID 属于另一个时间轮
467 ///
468 /// # 注意
469 /// - 任务 ID 在延期后保持不变
470 /// - 原始超时通知保持有效
471 /// - 如果 callback 为 `Some`,将替换原始回调
472 /// - 如果 callback 为 `None`,保留原始回调
473 ///
474 /// # Examples (示例)
475 /// ```no_run
476 /// # use kestrel_timer::{TimerWheel, TimerService, CallbackWrapper, TimerTask, config::ServiceConfig};
477 /// # use std::time::Duration;
478 /// #
479 /// # #[tokio::main]
480 /// # async fn main() {
481 /// let timer = TimerWheel::with_defaults();
482 /// let service = timer.create_service(ServiceConfig::default());
483 ///
484 /// let handle = service.allocate_handle();
485 /// let task_id = handle.task_id();
486 /// let callback = Some(CallbackWrapper::new(|| async {
487 /// println!("Original callback"); // 原始回调
488 /// }));
489 /// let task = TimerTask::new_oneshot(Duration::from_secs(5), callback);
490 /// service.register(handle, task).unwrap(); // 注册定时器
491 ///
492 /// // Postpone and replace callback (延期并替换回调)
493 /// let new_callback = Some(CallbackWrapper::new(|| async { println!("New callback!"); }));
494 /// let success = service
495 /// .postpone(task_id, Duration::from_secs(10), new_callback)
496 /// .unwrap();
497 /// println!("Postponed successfully: {}", success);
498 /// # }
499 /// ```
500 #[inline]
501 pub fn postpone(
502 &self,
503 task_id: TaskId,
504 new_delay: Duration,
505 callback: Option<CallbackWrapper>,
506 ) -> Result<bool, TimerError> {
507 let mut wheel = self.wheel.lock();
508 wheel.postpone_at(task_id, new_delay, callback)
509 }
510
511 /// Batch postpone tasks (keep original callbacks)
512 ///
513 /// # Parameters
514 /// - `updates`: List of tuples of (task ID, new delay)
515 ///
516 /// # Returns
517 /// Number of successfully postponed tasks, or
518 /// `Err(TimerError::WrongWheel)` if any ID belongs to another wheel.
519 ///
520 /// 批量延期任务 (保持原始回调)
521 ///
522 /// # 参数
523 /// - `updates`: (任务 ID, 新延迟) 元组列表
524 ///
525 /// # 返回值
526 /// 成功延期的任务数量;如果任一 ID 属于其他时间轮则返回
527 /// `Err(TimerError::WrongWheel)`,且不修改任何任务。
528 ///
529 /// # Examples (示例)
530 /// ```no_run
531 /// # use kestrel_timer::{TimerWheel, TimerService, CallbackWrapper, TimerTask, config::ServiceConfig};
532 /// # use std::time::Duration;
533 /// #
534 /// # #[tokio::main]
535 /// # async fn main() {
536 /// let timer = TimerWheel::with_defaults();
537 /// let service = timer.create_service(ServiceConfig::default());
538 ///
539 /// let handles = service.allocate_handles(3);
540 /// let task_ids: Vec<_> = handles.iter().map(|h| h.task_id()).collect();
541 /// let tasks: Vec<_> = (0..3)
542 /// .map(|i| {
543 /// let callback = Some(CallbackWrapper::new(move || async move {
544 /// println!("Timer {} fired!", i);
545 /// }));
546 /// TimerTask::new_oneshot(Duration::from_secs(5), callback)
547 /// })
548 /// .collect();
549 /// service.register_batch(handles, tasks).unwrap();
550 ///
551 /// // Batch postpone (keep original callbacks)
552 /// // 批量延期任务 (保持原始回调)
553 /// let updates: Vec<_> = task_ids
554 /// .into_iter()
555 /// .map(|id| (id, Duration::from_secs(10)))
556 /// .collect();
557 /// let postponed = service.postpone_batch(updates).unwrap();
558 /// println!("Postponed {} tasks", postponed);
559 /// # }
560 /// ```
561 #[inline]
562 pub fn postpone_batch(&self, updates: Vec<(TaskId, Duration)>) -> Result<usize, TimerError> {
563 if updates.is_empty() {
564 return Ok(0);
565 }
566
567 let mut wheel = self.wheel.lock();
568 wheel.postpone_batch_at(updates)
569 }
570
571 /// Batch postpone tasks (replace callbacks)
572 ///
573 /// # Parameters
574 /// - `updates`: List of tuples of (task ID, new delay, new callback)
575 ///
576 /// # Returns
577 /// Number of successfully postponed tasks, or
578 /// `Err(TimerError::WrongWheel)` if any ID belongs to another wheel.
579 ///
580 /// 批量延期任务 (替换回调)
581 ///
582 /// # 参数
583 /// - `updates`: (任务 ID, 新延迟, 新回调) 元组列表
584 ///
585 /// # 返回值
586 /// 成功延期的任务数量;如果任一 ID 属于其他时间轮则返回
587 /// `Err(TimerError::WrongWheel)`,且不修改任何任务。
588 ///
589 /// # Examples (示例)
590 /// ```no_run
591 /// # use kestrel_timer::{TimerWheel, TimerService, CallbackWrapper, config::ServiceConfig};
592 /// # use std::time::Duration;
593 /// #
594 /// # #[tokio::main]
595 /// # async fn main() {
596 /// # use kestrel_timer::TimerTask;
597 /// let timer = TimerWheel::with_defaults();
598 /// let service = timer.create_service(ServiceConfig::default());
599 ///
600 /// // Create 3 tasks, initially no callbacks
601 /// // 创建 3 个任务,最初没有回调
602 /// let handles = service.allocate_handles(3);
603 /// let task_ids: Vec<_> = handles.iter().map(|h| h.task_id()).collect();
604 /// let tasks: Vec<_> = (0..3)
605 /// .map(|_| {
606 /// TimerTask::new_oneshot(Duration::from_secs(5), None)
607 /// })
608 /// .collect();
609 /// service.register_batch(handles, tasks).unwrap();
610 ///
611 /// // Batch postpone and add new callbacks
612 /// // 批量延期并添加新的回调
613 /// let updates: Vec<_> = task_ids
614 /// .into_iter()
615 /// .enumerate()
616 /// .map(|(i, id)| {
617 /// let callback = Some(CallbackWrapper::new(move || async move {
618 /// println!("New callback {}", i);
619 /// }));
620 /// (id, Duration::from_secs(10), callback)
621 /// })
622 /// .collect();
623 /// let postponed = service.postpone_batch_with_callbacks(updates).unwrap();
624 /// println!("Postponed {} tasks", postponed);
625 /// # }
626 /// ```
627 #[inline]
628 pub fn postpone_batch_with_callbacks(
629 &self,
630 updates: Vec<(TaskId, Duration, Option<CallbackWrapper>)>,
631 ) -> Result<usize, TimerError> {
632 if updates.is_empty() {
633 return Ok(0);
634 }
635
636 let mut wheel = self.wheel.lock();
637 wheel.postpone_batch_with_callbacks_at(updates)
638 }
639
640 /// Register timer task to service (registration phase)
641 ///
642 /// # Parameters
643 /// - `handle`: Handle allocated via `allocate_handle()`
644 /// - `task`: Task created via `TimerTask::new_oneshot()` or `TimerTask::new_periodic()`
645 ///
646 /// # Returns
647 /// - `Ok(TimerHandle)`: Register successfully
648 /// - `Err(TimerError::RegisterFailed)`: Register failed (internal channel is full or closed)
649 /// - `Err(TimerError::WrongWheel)`: Handle belongs to another wheel
650 /// - `Err(TimerError::Shutdown)`: The timing wheel is closed
651 ///
652 /// 注册定时器任务到服务 (注册阶段)
653 /// # 参数
654 /// - `handle`: 通过 `allocate_handle()` 分配的 handle
655 /// - `task`: 通过 `TimerTask::new_oneshot()` 或 `TimerTask::new_periodic()` 创建的任务
656 ///
657 /// # 返回值
658 /// - `Ok(TimerHandle)`: 注册成功
659 /// - `Err(TimerError::RegisterFailed)`: 注册失败 (内部通道已满或关闭)
660 /// - `Err(TimerError::WrongWheel)`: handle 属于其他时间轮
661 /// - `Err(TimerError::Shutdown)`: 时间轮已关闭
662 ///
663 /// # Examples (示例)
664 /// ```no_run
665 /// # use kestrel_timer::{TimerWheel, TimerService, CallbackWrapper, config::ServiceConfig, TimerTask};
666 /// # use std::time::Duration;
667 /// #
668 /// # #[tokio::main]
669 /// # async fn main() {
670 /// let timer = TimerWheel::with_defaults();
671 /// let service = timer.create_service(ServiceConfig::default());
672 ///
673 /// // Step 1: allocate handle
674 /// // 分配 handle
675 /// let handle = service.allocate_handle();
676 /// let task_id = handle.task_id();
677 ///
678 /// // Step 2: create task
679 /// // 创建任务
680 /// let callback = Some(CallbackWrapper::new(|| async move {
681 /// println!("Timer fired!");
682 /// }));
683 /// let task = TimerTask::new_oneshot(Duration::from_millis(100), callback);
684 ///
685 /// // Step 3: register task
686 /// // 注册任务
687 /// service.register(handle, task).unwrap();
688 /// # }
689 /// ```
690 #[inline]
691 pub fn register(&self, handle: TaskHandle, task: TimerTask) -> Result<TimerHandle, TimerError> {
692 let task_id = handle.task_id();
693
694 let (task, completion_rx) = TimerTaskWithCompletionNotifier::from_timer_task(task);
695
696 // Single lock, complete all operations
697 // 单次锁定,完成所有操作
698 {
699 let mut wheel_guard = self.wheel.lock();
700 wheel_guard.insert_at(handle, task)?;
701 }
702
703 self.registered_task_ids.lock().insert(task_id);
704
705 // The wheel insertion is provisional until the actor accepts this command.
706 // 时间轮插入在 actor 接受命令前只是暂存状态。
707 match self.command_tx.try_send(ServiceCommand::AddTimerHandle {
708 task_id,
709 completion_rx,
710 }) {
711 Ok(()) => Ok(TimerHandle::new(task_id, self.wheel.clone())),
712 Err(error) => {
713 drop(error);
714 self.rollback_registration(&[task_id]);
715 Err(TimerError::RegisterFailed)
716 }
717 }
718 }
719
720 /// Batch register timer tasks to service (registration phase)
721 ///
722 /// # Parameters
723 /// - `handles`: Pre-allocated handles for tasks
724 /// - `tasks`: List of timer tasks
725 ///
726 /// # Returns
727 /// - `Ok(BatchHandle)`: Register successfully
728 /// - `Err(TimerError::RegisterFailed)`: Register failed (internal channel is full or closed)
729 /// - `Err(TimerError::BatchLengthMismatch)`: Handles and tasks lengths don't match
730 /// - `Err(TimerError::WrongWheel)`: Any handle belongs to another wheel
731 /// - `Err(TimerError::Shutdown)`: The timing wheel is closed
732 ///
733 /// 批量注册定时器任务到服务 (注册阶段)
734 /// # 参数
735 /// - `handles`: 任务的预分配 handles
736 /// - `tasks`: 定时器任务列表
737 ///
738 /// # 返回值
739 /// - `Ok(BatchHandle)`: 注册成功
740 /// - `Err(TimerError::RegisterFailed)`: 注册失败 (内部通道已满或关闭)
741 /// - `Err(TimerError::BatchLengthMismatch)`: handles 和 tasks 长度不匹配
742 /// - `Err(TimerError::WrongWheel)`: 任一 handle 属于其他时间轮
743 /// - `Err(TimerError::Shutdown)`: 时间轮已关闭
744 ///
745 /// # Examples (示例)
746 /// ```no_run
747 /// # use kestrel_timer::{TimerWheel, TimerService, CallbackWrapper, config::ServiceConfig, TimerTask};
748 /// # use std::time::Duration;
749 /// #
750 /// # #[tokio::main]
751 /// # async fn main() {
752 /// # use kestrel_timer::TimerTask;
753 /// let timer = TimerWheel::with_defaults();
754 /// let service = timer.create_service(ServiceConfig::default());
755 ///
756 /// // Step 1: batch allocate handles
757 /// // 批量分配 handles
758 /// let handles = service.allocate_handles(3);
759 ///
760 /// // Step 2: create tasks
761 /// // 创建任务
762 /// let tasks: Vec<_> = (0..3)
763 /// .map(|i| {
764 /// let callback = Some(CallbackWrapper::new(move || async move {
765 /// println!("Timer {} fired!", i);
766 /// }));
767 /// TimerTask::new_oneshot(Duration::from_secs(1), callback)
768 /// })
769 /// .collect();
770 ///
771 /// // Step 3: register batch
772 /// // 注册批量任务
773 /// service.register_batch(handles, tasks).unwrap();
774 /// # }
775 /// ```
776 #[inline]
777 pub fn register_batch(
778 &self,
779 handles: Vec<TaskHandle>,
780 tasks: Vec<TimerTask>,
781 ) -> Result<BatchHandle, TimerError> {
782 // Validate lengths match
783 if handles.len() != tasks.len() {
784 return Err(TimerError::BatchLengthMismatch {
785 handles_len: handles.len(),
786 tasks_len: tasks.len(),
787 });
788 }
789
790 let task_count = tasks.len();
791 let mut completion_rxs = Vec::with_capacity(task_count);
792 let mut task_ids = Vec::with_capacity(task_count);
793 let mut prepared_handles = Vec::with_capacity(task_count);
794 let mut prepared_tasks = Vec::with_capacity(task_count);
795
796 // Step 1: prepare all channels and notifiers (no lock)
797 // 步骤 1: 准备所有通道和通知器(无锁)
798 for (handle, task) in handles.into_iter().zip(tasks) {
799 let task_id = handle.task_id();
800 let (task, completion_rx) = TimerTaskWithCompletionNotifier::from_timer_task(task);
801 task_ids.push(task_id);
802 completion_rxs.push(completion_rx);
803 prepared_handles.push(handle);
804 prepared_tasks.push(task);
805 }
806
807 // Step 2: single lock, batch insert
808 // 步骤 2: 单次锁定,批量插入
809 {
810 let mut wheel_guard = self.wheel.lock();
811 wheel_guard.insert_batch_at(prepared_handles, prepared_tasks)?;
812 }
813
814 self.registered_task_ids
815 .lock()
816 .extend(task_ids.iter().copied());
817
818 // The batch insertion is provisional until the actor accepts this command.
819 // 批量插入在 actor 接受命令前只是暂存状态。
820 match self.command_tx.try_send(ServiceCommand::AddBatchHandle {
821 task_ids: task_ids.clone(),
822 completion_rxs,
823 }) {
824 Ok(()) => Ok(BatchHandle::new(task_ids, self.wheel.clone())),
825 Err(error) => {
826 drop(error);
827 self.rollback_registration(&task_ids);
828 Err(TimerError::RegisterFailed)
829 }
830 }
831 }
832
833 fn rollback_registration(&self, task_ids: &[TaskId]) {
834 if task_ids.is_empty() {
835 return;
836 }
837
838 self.registered_task_ids
839 .lock()
840 .retain(|task_id| !task_ids.contains(task_id));
841
842 let mut wheel = self.wheel.lock();
843 let _ = wheel.cancel_batch(task_ids);
844 }
845
846 fn cancel_registered_tasks(&self) {
847 let task_ids: Vec<_> = self.registered_task_ids.lock().drain().collect();
848
849 if task_ids.is_empty() {
850 return;
851 }
852
853 let mut wheel = self.wheel.lock();
854 let _ = wheel.cancel_batch(&task_ids);
855 }
856
857 /// Graceful shutdown of TimerService
858 ///
859 /// Cancels tasks registered through this service, closes its aggregated
860 /// notification receiver, and waits for the actor to stop. Tasks owned by
861 /// other services sharing the same wheel are left untouched.
862 ///
863 /// 优雅关闭 TimerService
864 ///
865 /// 取消通过此服务注册的任务,关闭聚合通知接收器,并等待 actor 停止。
866 /// 共享同一时间轮的其他服务任务不会受到影响。
867 ///
868 /// # Examples (示例)
869 /// ```no_run
870 /// # use kestrel_timer::{TimerWheel, config::ServiceConfig};
871 /// # #[tokio::main]
872 /// # async fn main() {
873 /// let timer = TimerWheel::with_defaults();
874 /// let mut service = timer.create_service(ServiceConfig::default());
875 ///
876 /// // Use service... (使用服务...)
877 ///
878 /// service.shutdown().await;
879 /// # }
880 /// ```
881 pub async fn shutdown(mut self) {
882 self.cancel_registered_tasks();
883
884 if let Some(shutdown_tx) = self.shutdown_tx.take() {
885 let _ = shutdown_tx.send(());
886 }
887 if let Some(handle) = self.actor_handle.take() {
888 let _ = handle.await;
889 }
890 }
891}
892
893impl Drop for TimerService {
894 fn drop(&mut self) {
895 if let Some(handle) = self.actor_handle.take() {
896 handle.abort();
897 }
898 }
899}
900
901/// ServiceActor - internal Actor implementation
902///
903/// ServiceActor - 内部 Actor 实现
904struct ServiceActor {
905 /// Command receiver
906 ///
907 /// 命令接收器
908 command_rx: mpsc::Receiver<ServiceCommand>,
909 /// Timeout sender (supports both one-shot and periodic task notifications)
910 ///
911 /// 超时发送器(支持一次性和周期性任务通知)
912 timeout_tx: spsc::Sender<TaskNotification, 32>,
913 /// Actor shutdown signal receiver
914 ///
915 /// Actor 关闭信号接收器
916 shutdown_rx: Receiver<()>,
917 /// Timing wheel lifecycle signal
918 ///
919 /// 时间轮生命周期信号
920 wheel_closed: watch::Receiver<bool>,
921 /// Shared registry used to remove completed service-owned tasks
922 ///
923 /// 用于移除已完成服务任务的共享注册表
924 registered_task_ids: Arc<Mutex<HashSet<TaskId>>>,
925}
926
927impl ServiceActor {
928 /// Create new ServiceActor
929 ///
930 /// 创建新的 ServiceActor
931 fn new(
932 command_rx: mpsc::Receiver<ServiceCommand>,
933 timeout_tx: spsc::Sender<TaskNotification, 32>,
934 shutdown_rx: Receiver<()>,
935 wheel_closed: watch::Receiver<bool>,
936 registered_task_ids: Arc<Mutex<HashSet<TaskId>>>,
937 ) -> Self {
938 Self {
939 command_rx,
940 timeout_tx,
941 shutdown_rx,
942 wheel_closed,
943 registered_task_ids,
944 }
945 }
946
947 async fn send_notification(
948 timeout_tx: &spsc::Sender<TaskNotification, 32>,
949 notification: TaskNotification,
950 shutdown_rx: &mut Receiver<()>,
951 wheel_closed: &mut watch::Receiver<bool>,
952 ) -> bool {
953 tokio::select! {
954 biased;
955 _ = &mut *shutdown_rx => false,
956 _ = wheel_closed.changed() => false,
957 result = timeout_tx.send(notification) => result.is_ok(),
958 }
959 }
960
961 /// Run Actor event loop
962 ///
963 /// 运行 Actor 事件循环
964 async fn run(mut self) {
965 // Use separate FuturesUnordered for one-shot and periodic tasks
966 // 为一次性任务和周期性任务使用独立的 FuturesUnordered
967
968 // One-shot futures: each future returns (TaskId, TaskCompletion)
969 // 一次性任务futures:每个future 返回 (TaskId, TaskCompletion)
970 let mut oneshot_futures: FuturesUnordered<BoxFuture<'static, (TaskId, TaskCompletion)>> =
971 FuturesUnordered::new();
972
973 // Periodic futures: each future returns (TaskId, Option<PeriodicTaskCompletion>, mpsc::Receiver)
974 // The receiver is returned so we can continue listening for next event
975 // 周期性任务futures:每个future 返回 (TaskId, Option<PeriodicTaskCompletion>, mpsc::Receiver)
976 // 返回接收器以便我们可以继续监听下一个事件
977 type PeriodicFutureResult = (
978 TaskId,
979 Option<TaskCompletion>,
980 crate::task::PeriodicCompletionReceiver,
981 );
982 let mut periodic_futures: FuturesUnordered<BoxFuture<'static, PeriodicFutureResult>> =
983 FuturesUnordered::new();
984
985 // Move shutdown_rx out of self, so it can be used in select! with &mut
986 // 将 shutdown_rx 从 self 中移出,以便在 select! 中使用 &mut
987 let timeout_tx = self.timeout_tx;
988 let mut shutdown_rx = self.shutdown_rx;
989 let mut wheel_closed = self.wheel_closed;
990 let registered_task_ids = self.registered_task_ids;
991
992 if *wheel_closed.borrow() {
993 return;
994 }
995
996 loop {
997 tokio::select! {
998 // Listen to high-priority shutdown signal
999 // 监听高优先级关闭信号
1000 _ = &mut shutdown_rx => {
1001 // Receive shutdown signal, exit loop immediately
1002 // 接收到关闭信号,立即退出循环
1003 break;
1004 }
1005
1006 // Stop when the owner timing wheel is shut down
1007 // 所属时间轮关闭时停止
1008 _ = wheel_closed.changed() => {
1009 break;
1010 }
1011
1012 // Listen to one-shot task timeout events
1013 // 监听一次性任务超时事件
1014 Some((task_id, completion)) = oneshot_futures.next() => {
1015 registered_task_ids.lock().remove(&task_id);
1016
1017 // Check completion reason, only forward Called events, do not forward Cancelled events
1018 // 检查完成原因,只转发 Called 事件,不转发 Cancelled 事件
1019 if completion == TaskCompletion::Called
1020 && !Self::send_notification(
1021 &timeout_tx,
1022 TaskNotification::OneShot(task_id),
1023 &mut shutdown_rx,
1024 &mut wheel_closed,
1025 )
1026 .await
1027 {
1028 break;
1029 }
1030 // Task will be automatically removed from FuturesUnordered
1031 // 任务将自动从 FuturesUnordered 中移除
1032 }
1033
1034 // Listen to periodic task events
1035 // 监听周期性任务事件
1036 Some((task_id, reason, mut receiver)) = periodic_futures.next() => {
1037 // Check completion reason, only forward Called events, do not forward Cancelled events
1038 // 检查完成原因,只转发 Called 事件,不转发 Cancelled 事件
1039 if let Some(TaskCompletion::Called) = reason {
1040 if !Self::send_notification(
1041 &timeout_tx,
1042 TaskNotification::Periodic(task_id),
1043 &mut shutdown_rx,
1044 &mut wheel_closed,
1045 )
1046 .await
1047 {
1048 break;
1049 }
1050
1051 // Re-add the receiver to continue listening for next periodic event
1052 // 重新添加接收器以继续监听下一个周期性事件
1053 let future: BoxFuture<'static, PeriodicFutureResult> = Box::pin(async move {
1054 let reason = receiver.recv().await;
1055 (task_id, reason, receiver)
1056 });
1057 periodic_futures.push(future);
1058 } else {
1059 registered_task_ids.lock().remove(&task_id);
1060 }
1061 // If Cancelled or None, do not re-add the future (task is done)
1062 // 如果 Cancelled 或 None,不重新添加 future(任务结束)
1063 }
1064
1065 // Listen to commands
1066 // 监听命令
1067 Some(cmd) = self.command_rx.recv() => {
1068 match cmd {
1069 ServiceCommand::AddBatchHandle { task_ids, completion_rxs } => {
1070 // Add all tasks to appropriate futures
1071 // 将所有任务添加到相应 futures
1072 for (task_id, rx) in task_ids.into_iter().zip(completion_rxs.into_iter()) {
1073 match rx {
1074 crate::task::CompletionReceiver::OneShot(receiver) => {
1075 let future: BoxFuture<'static, (TaskId, TaskCompletion)> = Box::pin(async move {
1076 // A dropped task closes the receiver without a completion event.
1077 // 任务被丢弃时接收器会关闭,此时没有完成事件。
1078 let completion = receiver
1079 .recv()
1080 .await
1081 .unwrap_or(TaskCompletion::Cancelled);
1082 (task_id, completion)
1083 });
1084 oneshot_futures.push(future);
1085 },
1086 crate::task::CompletionReceiver::Periodic(mut receiver) => {
1087 let future: BoxFuture<'static, PeriodicFutureResult> = Box::pin(async move {
1088 let reason = receiver.recv().await;
1089 (task_id, reason, receiver)
1090 });
1091 periodic_futures.push(future);
1092 }
1093 }
1094 }
1095 }
1096 ServiceCommand::AddTimerHandle { task_id, completion_rx } => {
1097 // Add to appropriate futures
1098 // 添加到相应的 futures
1099 match completion_rx {
1100 crate::task::CompletionReceiver::OneShot(receiver) => {
1101 let future: BoxFuture<'static, (TaskId, TaskCompletion)> = Box::pin(async move {
1102 // A dropped task closes the receiver without a completion event.
1103 // 任务被丢弃时接收器会关闭,此时没有完成事件。
1104 let completion = receiver
1105 .recv()
1106 .await
1107 .unwrap_or(TaskCompletion::Cancelled);
1108 (task_id, completion)
1109 });
1110 oneshot_futures.push(future);
1111 },
1112 crate::task::CompletionReceiver::Periodic(mut receiver) => {
1113 let future: BoxFuture<'static, PeriodicFutureResult> = Box::pin(async move {
1114 let reason = receiver.recv().await;
1115 (task_id, reason, receiver)
1116 });
1117 periodic_futures.push(future);
1118 }
1119 }
1120 }
1121 }
1122 }
1123
1124 // If no futures and command channel is closed, exit loop
1125 // 如果没有 futures 且命令通道关闭,退出循环
1126 else => {
1127 break;
1128 }
1129 }
1130 }
1131 }
1132}
1133
1134#[cfg(test)]
1135mod tests {
1136 use super::*;
1137 use crate::{TimerTask, TimerWheel};
1138 use std::num::NonZeroUsize;
1139 use std::sync::Arc;
1140 use std::sync::atomic::{AtomicU32, Ordering};
1141 use std::time::Duration;
1142
1143 #[tokio::test]
1144 async fn test_service_creation() {
1145 let timer = TimerWheel::with_defaults();
1146 let _service = timer.create_service(ServiceConfig::default());
1147 }
1148
1149 #[tokio::test]
1150 async fn test_add_timer_handle_and_receive_timeout() {
1151 let timer = TimerWheel::with_defaults();
1152 let mut service = timer.create_service(ServiceConfig::default());
1153
1154 // Allocate handle (分配 handle)
1155 let handle = service.allocate_handle();
1156 let task_id = handle.task_id();
1157
1158 // Create single timer (创建单个定时器)
1159 let task = TimerTask::new_oneshot(
1160 Duration::from_millis(50),
1161 Some(CallbackWrapper::new(|| async {})),
1162 );
1163
1164 // Register to service (注册到服务)
1165 service.register(handle, task).unwrap();
1166
1167 // Receive timeout notification (接收超时通知)
1168 let rx = service.take_receiver().unwrap();
1169 let received_notification = tokio::time::timeout(Duration::from_millis(200), rx.recv())
1170 .await
1171 .expect("Should receive timeout notification")
1172 .expect("Should receive Some value");
1173
1174 assert_eq!(received_notification, TaskNotification::OneShot(task_id));
1175 }
1176
1177 #[tokio::test]
1178 async fn test_shutdown() {
1179 let timer = TimerWheel::with_defaults();
1180 let service = timer.create_service(ServiceConfig::default());
1181
1182 // Add some timers (添加一些定时器)
1183 let handle1 = service.allocate_handle();
1184 let handle2 = service.allocate_handle();
1185 let task1 = TimerTask::new_oneshot(Duration::from_secs(10), None);
1186 let task2 = TimerTask::new_oneshot(Duration::from_secs(10), None);
1187 service.register(handle1, task1).unwrap();
1188 service.register(handle2, task2).unwrap();
1189
1190 // Immediately shutdown (without waiting for timers to trigger) (立即关闭(不等待定时器触发))
1191 service.shutdown().await;
1192 }
1193
1194 #[tokio::test]
1195 async fn test_shutdown_cancels_owned_tasks_and_closes_receiver() {
1196 let timer = TimerWheel::with_defaults();
1197 let mut service = timer.create_service(ServiceConfig::default());
1198 let callback_count = Arc::new(AtomicU32::new(0));
1199 let callback_count_clone = Arc::clone(&callback_count);
1200 let handle = service
1201 .register(
1202 service.allocate_handle(),
1203 TimerTask::new_oneshot(
1204 Duration::from_secs(10),
1205 Some(CallbackWrapper::new(move || {
1206 let callback_count = Arc::clone(&callback_count_clone);
1207 async move {
1208 callback_count.fetch_add(1, Ordering::SeqCst);
1209 }
1210 })),
1211 ),
1212 )
1213 .unwrap();
1214 let receiver = service.take_receiver().unwrap();
1215
1216 service.shutdown().await;
1217
1218 assert!(
1219 tokio::time::timeout(Duration::from_millis(100), receiver.recv())
1220 .await
1221 .expect("service receiver should close")
1222 .is_none()
1223 );
1224 assert!(!handle.cancel().unwrap());
1225 assert_eq!(callback_count.load(Ordering::SeqCst), 0);
1226 }
1227
1228 #[tokio::test]
1229 async fn test_timer_shutdown_closes_bound_service() {
1230 let timer = TimerWheel::with_defaults();
1231 let mut service = timer.create_service(ServiceConfig::default());
1232 let receiver = service.take_receiver().unwrap();
1233
1234 timer.shutdown().await;
1235
1236 assert!(
1237 tokio::time::timeout(Duration::from_millis(100), receiver.recv())
1238 .await
1239 .expect("service receiver should close when its wheel closes")
1240 .is_none()
1241 );
1242 assert!(matches!(
1243 service.register(
1244 service.allocate_handle(),
1245 TimerTask::new_oneshot(Duration::from_secs(1), None),
1246 ),
1247 Err(TimerError::Shutdown)
1248 ));
1249 service.shutdown().await;
1250 }
1251
1252 #[tokio::test(flavor = "current_thread")]
1253 async fn test_shutdown_when_timeout_channel_is_full() {
1254 let config = ServiceConfig::builder()
1255 .timeout_channel_capacity(NonZeroUsize::new(1).unwrap())
1256 .build();
1257 let timer = TimerWheel::with_defaults();
1258 let mut service = timer.create_service(config);
1259
1260 let first_handle = service.allocate_handle();
1261 service
1262 .register(
1263 first_handle,
1264 TimerTask::new_oneshot(Duration::from_millis(10), None),
1265 )
1266 .unwrap();
1267
1268 let second_handle = service.allocate_handle();
1269 service
1270 .register(
1271 second_handle,
1272 TimerTask::new_oneshot(Duration::from_millis(20), None),
1273 )
1274 .unwrap();
1275
1276 // Keep the receiver alive without consuming it so the second notification
1277 // fills the output channel and blocks the actor's normal send path.
1278 let _receiver = service.take_receiver().unwrap();
1279 tokio::time::sleep(Duration::from_millis(80)).await;
1280
1281 tokio::time::timeout(Duration::from_millis(100), service.shutdown())
1282 .await
1283 .expect("shutdown should not wait for a full timeout channel");
1284 }
1285
1286 #[tokio::test]
1287 async fn test_schedule_once_direct() {
1288 let timer = TimerWheel::with_defaults();
1289 let mut service = timer.create_service(ServiceConfig::default());
1290 let counter = Arc::new(AtomicU32::new(0));
1291
1292 // Schedule timer directly through service
1293 // 直接通过服务调度定时器
1294 let counter_clone = Arc::clone(&counter);
1295 let handle = service.allocate_handle();
1296 let task_id = handle.task_id();
1297 let task = TimerTask::new_oneshot(
1298 Duration::from_millis(50),
1299 Some(CallbackWrapper::new(move || {
1300 let counter = Arc::clone(&counter_clone);
1301 async move {
1302 counter.fetch_add(1, Ordering::SeqCst);
1303 }
1304 })),
1305 );
1306 service.register(handle, task).unwrap();
1307
1308 // Wait for timer to trigger
1309 // 等待定时器触发
1310 let rx = service.take_receiver().unwrap();
1311 let received_notification = tokio::time::timeout(Duration::from_millis(200), rx.recv())
1312 .await
1313 .expect("Should receive timeout notification")
1314 .expect("Should receive Some value");
1315
1316 assert_eq!(received_notification, TaskNotification::OneShot(task_id));
1317
1318 // Wait for callback to execute
1319 // 等待回调执行
1320 tokio::time::sleep(Duration::from_millis(50)).await;
1321 assert_eq!(counter.load(Ordering::SeqCst), 1);
1322 }
1323
1324 #[tokio::test]
1325 async fn test_schedule_once_notify_direct() {
1326 let timer = TimerWheel::with_defaults();
1327 let mut service = timer.create_service(ServiceConfig::default());
1328
1329 // Schedule only notification timer directly through service (no callback)
1330 // 直接通过服务调度通知定时器(没有回调函数)
1331 let handle = service.allocate_handle();
1332 let task_id = handle.task_id();
1333 let task = TimerTask::new_oneshot(Duration::from_millis(50), None);
1334 service.register(handle, task).unwrap();
1335
1336 // Receive timeout notification
1337 // 接收超时通知
1338 let rx = service.take_receiver().unwrap();
1339 let received_notification = tokio::time::timeout(Duration::from_millis(200), rx.recv())
1340 .await
1341 .expect("Should receive timeout notification")
1342 .expect("Should receive Some value");
1343
1344 assert_eq!(received_notification, TaskNotification::OneShot(task_id));
1345 }
1346
1347 #[tokio::test]
1348 async fn test_task_timeout_cleans_up_task_sender() {
1349 let timer = TimerWheel::with_defaults();
1350 let mut service = timer.create_service(ServiceConfig::default());
1351
1352 // Add a short-term timer (添加短期定时器)
1353 let handle = service.allocate_handle();
1354 let task_id = handle.task_id();
1355 let task = TimerTask::new_oneshot(Duration::from_millis(50), None);
1356
1357 service.register(handle, task).unwrap();
1358
1359 // Wait for task timeout (等待任务超时)
1360 let rx = service.take_receiver().unwrap();
1361 let received_notification = tokio::time::timeout(Duration::from_millis(200), rx.recv())
1362 .await
1363 .expect("Should receive timeout notification")
1364 .expect("Should receive Some value");
1365
1366 assert_eq!(received_notification, TaskNotification::OneShot(task_id));
1367
1368 // Wait a moment to ensure internal cleanup is complete (等待片刻以确保内部清理完成)
1369 tokio::time::sleep(Duration::from_millis(10)).await;
1370
1371 // Try to cancel the timed-out task, should return false (尝试取消超时任务,应返回 false)
1372 let cancelled = service.cancel_task(task_id).unwrap();
1373 assert!(!cancelled, "Timed out task should not exist anymore");
1374 }
1375
1376 #[tokio::test]
1377 async fn test_take_receiver_twice() {
1378 let timer = TimerWheel::with_defaults();
1379 let mut service = timer.create_service(ServiceConfig::default());
1380
1381 // First call should return Some
1382 // 第一次调用应该返回 Some
1383 let rx1 = service.take_receiver();
1384 assert!(rx1.is_some(), "First take_receiver should return Some");
1385
1386 // Second call should return None
1387 // 第二次调用应该返回 None
1388 let rx2 = service.take_receiver();
1389 assert!(rx2.is_none(), "Second take_receiver should return None");
1390 }
1391
1392 #[tokio::test]
1393 async fn test_concurrent_registration_single_service() {
1394 let timer = TimerWheel::with_defaults();
1395 let service = Arc::new(timer.create_service(ServiceConfig::default()));
1396
1397 let mut tasks = Vec::new();
1398 for _ in 0..10 {
1399 let service_clone = Arc::clone(&service);
1400 tasks.push(tokio::spawn(async move {
1401 for _ in 0..20 {
1402 let handle = service_clone.allocate_handle();
1403 let task = TimerTask::new_oneshot(Duration::from_millis(500), None);
1404 service_clone
1405 .register(handle, task)
1406 .expect("register should succeed");
1407 }
1408 }));
1409 }
1410
1411 for task in tasks {
1412 task.await.unwrap();
1413 }
1414 }
1415
1416 #[tokio::test(flavor = "current_thread")]
1417 async fn test_register_failure_rolls_back_single_task_when_channel_is_full() {
1418 let config = ServiceConfig::builder()
1419 .command_channel_capacity(NonZeroUsize::new(1).unwrap())
1420 .build();
1421 let timer = TimerWheel::with_defaults();
1422 let service = timer.create_service(config);
1423
1424 let first_handle = service.allocate_handle();
1425 let first_task_id = first_handle.task_id();
1426 service
1427 .register(
1428 first_handle,
1429 TimerTask::new_oneshot(Duration::from_secs(10), None),
1430 )
1431 .unwrap();
1432
1433 let failed_handle = service.allocate_handle();
1434 let failed_task_id = failed_handle.task_id();
1435 let callback_count = Arc::new(AtomicU32::new(0));
1436 let callback_count_clone = Arc::clone(&callback_count);
1437 let result = service.register(
1438 failed_handle,
1439 TimerTask::new_periodic(
1440 Duration::from_millis(10),
1441 Duration::from_millis(10),
1442 Some(CallbackWrapper::new(move || {
1443 let callback_count = Arc::clone(&callback_count_clone);
1444 async move {
1445 callback_count.fetch_add(1, Ordering::SeqCst);
1446 }
1447 })),
1448 None,
1449 )
1450 .unwrap(),
1451 );
1452
1453 assert!(matches!(result, Err(TimerError::RegisterFailed)));
1454 assert!(!service.cancel_task(failed_task_id).unwrap());
1455 assert!(service.cancel_task(first_task_id).unwrap());
1456 assert!(service.wheel.lock().is_empty());
1457 tokio::time::sleep(Duration::from_millis(50)).await;
1458 assert_eq!(callback_count.load(Ordering::SeqCst), 0);
1459 }
1460
1461 #[tokio::test(flavor = "current_thread")]
1462 async fn test_register_batch_failure_rolls_back_all_tasks_when_channel_is_full() {
1463 let config = ServiceConfig::builder()
1464 .command_channel_capacity(NonZeroUsize::new(1).unwrap())
1465 .build();
1466 let timer = TimerWheel::with_defaults();
1467 let service = timer.create_service(config);
1468
1469 let first_handle = service.allocate_handle();
1470 let first_task_id = first_handle.task_id();
1471 service
1472 .register(
1473 first_handle,
1474 TimerTask::new_oneshot(Duration::from_secs(10), None),
1475 )
1476 .unwrap();
1477
1478 let failed_handles = service.allocate_handles(2);
1479 let failed_task_ids: Vec<_> = failed_handles
1480 .iter()
1481 .map(|handle| handle.task_id())
1482 .collect();
1483 let failed_tasks = vec![
1484 TimerTask::new_oneshot(Duration::from_secs(10), None),
1485 TimerTask::new_oneshot(Duration::from_secs(10), None),
1486 ];
1487 let result = service.register_batch(failed_handles, failed_tasks);
1488
1489 assert!(matches!(result, Err(TimerError::RegisterFailed)));
1490 assert_eq!(service.cancel_batch(&failed_task_ids).unwrap(), 0);
1491 assert!(service.cancel_task(first_task_id).unwrap());
1492 assert!(service.wheel.lock().is_empty());
1493 }
1494
1495 #[tokio::test]
1496 async fn test_register_failure_rolls_back_when_actor_is_closed() {
1497 let timer = TimerWheel::with_defaults();
1498 let mut service = timer.create_service(ServiceConfig::default());
1499 let actor_handle = service.actor_handle.take().unwrap();
1500 actor_handle.abort();
1501 let _ = actor_handle.await;
1502
1503 let handle = service.allocate_handle();
1504 let task_id = handle.task_id();
1505 let result = service.register(
1506 handle,
1507 TimerTask::new_oneshot(Duration::from_secs(10), None),
1508 );
1509
1510 assert!(matches!(result, Err(TimerError::RegisterFailed)));
1511 assert!(!service.cancel_task(task_id).unwrap());
1512 assert!(service.wheel.lock().is_empty());
1513 }
1514
1515 #[tokio::test]
1516 async fn test_register_batch_failure_rolls_back_when_actor_is_closed() {
1517 let timer = TimerWheel::with_defaults();
1518 let mut service = timer.create_service(ServiceConfig::default());
1519 let actor_handle = service.actor_handle.take().unwrap();
1520 actor_handle.abort();
1521 let _ = actor_handle.await;
1522
1523 let handles = service.allocate_handles(2);
1524 let task_ids: Vec<_> = handles.iter().map(|handle| handle.task_id()).collect();
1525 let tasks = vec![
1526 TimerTask::new_oneshot(Duration::from_secs(10), None),
1527 TimerTask::new_oneshot(Duration::from_secs(10), None),
1528 ];
1529 let result = service.register_batch(handles, tasks);
1530
1531 assert!(matches!(result, Err(TimerError::RegisterFailed)));
1532 assert_eq!(service.cancel_batch(&task_ids).unwrap(), 0);
1533 assert!(service.wheel.lock().is_empty());
1534 }
1535}