Skip to main content

fynd_core/worker_pool/
task_queue.rs

1//! Task queue for distributing solve requests to workers.
2//!
3//! The queue sits between the HTTP handlers and the worker pool.
4//! It provides backpressure and allows the HTTP layer to remain
5//! responsive even when workers are busy.
6
7use std::time::Instant;
8
9use tokio::sync::oneshot;
10use uuid::Uuid;
11
12use crate::{types::internal::SolveTask, Order, SingleOrderQuote, SolveError, SolveParams};
13
14/// Configuration for the task queue.
15#[derive(Debug, Clone)]
16pub struct TaskQueueConfig {
17    /// Maximum number of pending tasks.
18    pub capacity: usize,
19}
20
21impl Default for TaskQueueConfig {
22    fn default() -> Self {
23        Self { capacity: 1000 }
24    }
25}
26
27/// Handle for enqueueing tasks.
28///
29/// This is cloned and shared with HTTP handlers.
30#[derive(Clone)]
31pub struct TaskQueueHandle {
32    sender: async_channel::Sender<SolveTask>,
33}
34
35impl TaskQueueHandle {
36    /// Enqueues a solve request and returns a future that resolves to the result.
37    ///
38    /// Returns an error if the queue is full.
39    pub async fn enqueue(
40        &self,
41        order: Order,
42        params: SolveParams,
43        deadline: Instant,
44    ) -> Result<SingleOrderQuote, SolveError> {
45        // Create response channel
46        let (response_tx, response_rx) = oneshot::channel();
47
48        // Generate task ID
49        let task_id = Uuid::new_v4();
50
51        // Create task
52        let task = SolveTask::new(task_id, order, response_tx, deadline).with_params(params);
53
54        // Try to send
55        self.sender
56            .send(task)
57            .await
58            .map_err(|_| SolveError::QueueFull)?;
59
60        // Wait for response
61        response_rx
62            .await
63            .map_err(|_| SolveError::Internal("worker dropped response channel".to_string()))?
64    }
65
66    /// Returns the current approximate queue depth.
67    ///
68    /// Note: This is not exact due to the async nature of the queue.
69    #[cfg(test)]
70    pub fn approximate_depth(&self) -> usize {
71        self.sender.len()
72    }
73
74    /// Returns true if the queue is likely full.
75    #[cfg(test)]
76    pub fn is_full(&self) -> bool {
77        self.sender.is_full()
78    }
79
80    /// Creates a TaskQueueHandle from an existing sender.
81    ///
82    /// This is primarily useful for testing with mock channels.
83    pub fn from_sender(sender: async_channel::Sender<SolveTask>) -> Self {
84        Self { sender }
85    }
86}
87
88/// The task queue itself.
89///
90/// This is consumed when creating the worker pool.
91pub struct TaskQueue {
92    receiver: async_channel::Receiver<SolveTask>,
93    handle: TaskQueueHandle,
94}
95
96impl TaskQueue {
97    /// Creates a new task queue with the given configuration.
98    pub fn new(config: TaskQueueConfig) -> Self {
99        let (sender, receiver) = async_channel::bounded(config.capacity);
100        let handle = TaskQueueHandle { sender };
101
102        Self { receiver, handle }
103    }
104
105    /// Splits the queue into handle and receiver.
106    pub fn split(self) -> (TaskQueueHandle, async_channel::Receiver<SolveTask>) {
107        (self.handle, self.receiver)
108    }
109
110    /// Returns a handle for enqueueing tasks.
111    #[cfg(test)]
112    pub fn handle(&self) -> TaskQueueHandle {
113        self.handle.clone()
114    }
115
116    /// Consumes the queue and returns the receiver.
117    ///
118    /// This is called when setting up the worker pool.
119    #[cfg(test)]
120    pub fn into_receiver(self) -> async_channel::Receiver<SolveTask> {
121        self.receiver
122    }
123}
124
125#[cfg(test)]
126mod tests {
127    use num_bigint::BigUint;
128    use rstest::rstest;
129    use tycho_simulation::tycho_core::{models::Address, Bytes};
130
131    use super::*;
132    use crate::{
133        BlockInfo, Order, OrderQuote, OrderSide, QuoteStatus, SingleOrderQuote, SolveParams,
134    };
135
136    // -------------------------------------------------------------------------
137    // Test Helpers
138    // -------------------------------------------------------------------------
139
140    fn make_address(byte: u8) -> Address {
141        Address::from([byte; 20])
142    }
143
144    /// A deadline long enough that these tests never reach it, so they measure the queue
145    /// rather than the router giving up.
146    fn long_deadline() -> Instant {
147        Instant::now() + std::time::Duration::from_secs(60)
148    }
149
150    fn make_order() -> Order {
151        Order::new(
152            make_address(0x01),
153            make_address(0x02),
154            BigUint::from(1000u64),
155            OrderSide::Sell,
156            make_address(0xAA),
157        )
158        .with_id("test-order".to_string())
159    }
160
161    fn make_single_quote() -> SingleOrderQuote {
162        SingleOrderQuote::new(
163            OrderQuote::new(
164                "test-order".to_string(),
165                QuoteStatus::Success,
166                BigUint::from(1000u64),
167                BigUint::from(990u64),
168                BigUint::from(100_000u64),
169                BigUint::from(990u64),
170                BlockInfo::new(1, "0x123".to_string(), 1000),
171                "test".to_string(),
172                Bytes::from(make_address(0xAA).as_ref()),
173                Bytes::from(make_address(0xAA).as_ref()),
174                "1".to_string(),
175            ),
176            5,
177        )
178    }
179
180    // -------------------------------------------------------------------------
181    // TaskQueueConfig Tests
182    // -------------------------------------------------------------------------
183
184    #[test]
185    fn test_config_default() {
186        let config = TaskQueueConfig::default();
187        assert_eq!(config.capacity, 1000);
188    }
189
190    #[rstest]
191    #[case::small(1)]
192    #[case::medium(100)]
193    #[case::large(10_000)]
194    fn test_config_custom_capacity(#[case] capacity: usize) {
195        let config = TaskQueueConfig { capacity };
196        assert_eq!(config.capacity, capacity);
197    }
198
199    // -------------------------------------------------------------------------
200    // TaskQueue Creation Tests
201    // -------------------------------------------------------------------------
202
203    #[rstest]
204    #[case::capacity_1(1)]
205    #[case::capacity_10(10)]
206    #[case::capacity_100(100)]
207    fn test_queue_creation(#[case] capacity: usize) {
208        let config = TaskQueueConfig { capacity };
209        let queue = TaskQueue::new(config);
210        let handle = queue.handle();
211
212        assert!(!handle.is_full());
213        assert_eq!(handle.approximate_depth(), 0);
214    }
215
216    #[test]
217    fn test_queue_handle_is_cloneable() {
218        let queue = TaskQueue::new(TaskQueueConfig { capacity: 10 });
219        let handle1 = queue.handle();
220        let handle2 = handle1.clone();
221
222        // Both handles should report same state
223        assert_eq!(handle1.approximate_depth(), handle2.approximate_depth());
224        assert_eq!(handle1.is_full(), handle2.is_full());
225    }
226
227    #[test]
228    fn test_queue_into_receiver() {
229        let queue = TaskQueue::new(TaskQueueConfig { capacity: 10 });
230        let _handle = queue.handle();
231        let _receiver = queue.into_receiver();
232        // Queue is consumed - receiver is ready for worker pool
233    }
234
235    #[test]
236    fn test_queue_split() {
237        let queue = TaskQueue::new(TaskQueueConfig { capacity: 10 });
238        let (handle, _receiver) = queue.split();
239
240        assert!(!handle.is_full());
241        assert_eq!(handle.approximate_depth(), 0);
242    }
243
244    // -------------------------------------------------------------------------
245    // TaskQueueHandle Tests
246    // -------------------------------------------------------------------------
247
248    #[tokio::test]
249    async fn test_enqueue_and_receive_response() {
250        let queue = TaskQueue::new(TaskQueueConfig { capacity: 10 });
251        let handle = queue.handle();
252        let receiver = queue.into_receiver();
253
254        // Spawn a "worker" that responds to the task
255        let worker = tokio::spawn(async move {
256            let task = receiver
257                .recv()
258                .await
259                .expect("should receive task");
260            assert_eq!(task.order().id(), "test-order");
261            task.respond(Ok(make_single_quote()));
262        });
263
264        // Enqueue an order
265        let result = handle
266            .enqueue(make_order(), SolveParams::default(), long_deadline())
267            .await;
268
269        worker
270            .await
271            .expect("worker should complete");
272        let quote = result.expect("should get quote");
273        assert_eq!(quote.solve_time_ms(), 5);
274    }
275
276    #[tokio::test]
277    async fn test_enqueue_receives_error_response() {
278        let queue = TaskQueue::new(TaskQueueConfig { capacity: 10 });
279        let handle = queue.handle();
280        let receiver = queue.into_receiver();
281
282        let worker = tokio::spawn(async move {
283            let task = receiver
284                .recv()
285                .await
286                .expect("should receive task");
287            task.respond(Err(SolveError::no_route_found("test")));
288        });
289
290        let result = handle
291            .enqueue(make_order(), SolveParams::default(), long_deadline())
292            .await;
293
294        worker
295            .await
296            .expect("worker should complete");
297        assert!(matches!(result, Err(SolveError::NoRouteFound { .. })));
298    }
299
300    #[tokio::test]
301    async fn test_enqueue_error_when_receiver_dropped() {
302        let queue = TaskQueue::new(TaskQueueConfig { capacity: 10 });
303        let handle = queue.handle();
304        let receiver = queue.into_receiver();
305
306        // Worker receives task but drops it without responding
307        let worker = tokio::spawn(async move {
308            let task = receiver
309                .recv()
310                .await
311                .expect("should receive task");
312            drop(task); // Drop without responding
313        });
314
315        let result = handle
316            .enqueue(make_order(), SolveParams::default(), long_deadline())
317            .await;
318
319        worker
320            .await
321            .expect("worker should complete");
322        assert!(matches!(result, Err(SolveError::Internal(_))));
323    }
324
325    #[tokio::test]
326    async fn test_enqueue_queue_full_error() {
327        let queue = TaskQueue::new(TaskQueueConfig { capacity: 10 });
328        let handle = queue.handle();
329        let receiver = queue.into_receiver();
330
331        // Drop receiver to close channel
332        drop(receiver);
333
334        let result = handle
335            .enqueue(make_order(), SolveParams::default(), long_deadline())
336            .await;
337        assert!(matches!(result, Err(SolveError::QueueFull)));
338    }
339
340    // -------------------------------------------------------------------------
341    // Queue Depth and Full Detection Tests
342    // -------------------------------------------------------------------------
343
344    #[tokio::test]
345    async fn test_approximate_depth_increases_with_pending_tasks() {
346        let queue = TaskQueue::new(TaskQueueConfig { capacity: 10 });
347        let handle = queue.handle();
348        let _receiver = queue.into_receiver(); // Keep receiver alive but don't consume
349
350        // Create a oneshot and send a task
351        let (response_tx, _response_rx) = oneshot::channel();
352        let task = SolveTask::new(Uuid::new_v4(), make_order(), response_tx, long_deadline());
353
354        handle
355            .sender
356            .send(task)
357            .await
358            .expect("should send");
359
360        assert_eq!(handle.approximate_depth(), 1);
361
362        // Send another
363        let (response_tx2, _response_rx2) = oneshot::channel();
364        let task2 = SolveTask::new(Uuid::new_v4(), make_order(), response_tx2, long_deadline());
365        handle
366            .sender
367            .send(task2)
368            .await
369            .expect("should send");
370
371        assert_eq!(handle.approximate_depth(), 2);
372    }
373
374    #[rstest]
375    #[case::capacity_1(1)]
376    #[case::capacity_5(5)]
377    #[case::capacity_10(10)]
378    #[tokio::test]
379    async fn test_is_full_when_at_capacity(#[case] capacity: usize) {
380        let queue = TaskQueue::new(TaskQueueConfig { capacity });
381        let handle = queue.handle();
382        let _receiver = queue.into_receiver();
383
384        // Fill the queue
385        for _ in 0..capacity {
386            let (response_tx, _response_rx) = oneshot::channel();
387            let task = SolveTask::new(Uuid::new_v4(), make_order(), response_tx, long_deadline());
388            handle
389                .sender
390                .send(task)
391                .await
392                .expect("should send");
393        }
394
395        assert!(handle.is_full());
396        assert_eq!(handle.approximate_depth(), capacity);
397    }
398
399    #[tokio::test]
400    async fn test_is_full_becomes_false_after_task_consumed() {
401        let queue = TaskQueue::new(TaskQueueConfig { capacity: 2 });
402        let handle = queue.handle();
403        let receiver = queue.into_receiver();
404
405        // Fill queue
406        let (tx1, _rx1) = oneshot::channel();
407        let (tx2, _rx2) = oneshot::channel();
408        handle
409            .sender
410            .send(SolveTask::new(Uuid::new_v4(), make_order(), tx1, long_deadline()))
411            .await
412            .unwrap();
413        handle
414            .sender
415            .send(SolveTask::new(Uuid::new_v4(), make_order(), tx2, long_deadline()))
416            .await
417            .unwrap();
418
419        assert!(handle.is_full());
420
421        // Consume one task
422        let _task = receiver.recv().await.unwrap();
423
424        // Queue should no longer be full
425        assert!(!handle.is_full());
426        assert_eq!(handle.approximate_depth(), 1);
427    }
428
429    // -------------------------------------------------------------------------
430    // Concurrent Operation Tests
431    // -------------------------------------------------------------------------
432
433    #[tokio::test]
434    async fn test_multiple_handles_can_enqueue_concurrently() {
435        let queue = TaskQueue::new(TaskQueueConfig { capacity: 10 });
436        let handle1 = queue.handle();
437        let handle2 = queue.handle();
438        let receiver = queue.into_receiver();
439
440        // Spawn worker that processes multiple tasks
441        let worker = tokio::spawn(async move {
442            for _ in 0..2 {
443                let task = receiver
444                    .recv()
445                    .await
446                    .expect("should receive task");
447                task.respond(Ok(make_single_quote()));
448            }
449        });
450
451        // Enqueue from both handles concurrently
452        let (result1, result2) = tokio::join!(
453            handle1.enqueue(make_order(), SolveParams::default(), long_deadline()),
454            handle2.enqueue(make_order(), SolveParams::default(), long_deadline()),
455        );
456
457        worker
458            .await
459            .expect("worker should complete");
460
461        assert!(result1.is_ok());
462        assert!(result2.is_ok());
463    }
464
465    #[tokio::test]
466    async fn test_task_id_is_unique_per_enqueue() {
467        let queue = TaskQueue::new(TaskQueueConfig { capacity: 10 });
468        let handle = queue.handle();
469        let receiver = queue.into_receiver();
470
471        // Spawn workers to collect task IDs
472        let collector = tokio::spawn(async move {
473            let task1 = receiver.recv().await.unwrap();
474            let id1 = task1.id();
475            task1.respond(Ok(make_single_quote()));
476
477            let task2 = receiver.recv().await.unwrap();
478            let id2 = task2.id();
479            task2.respond(Ok(make_single_quote()));
480
481            (id1, id2)
482        });
483
484        // Enqueue two orders
485        let _ = handle
486            .enqueue(make_order(), SolveParams::default(), long_deadline())
487            .await;
488        let _ = handle
489            .enqueue(make_order(), SolveParams::default(), long_deadline())
490            .await;
491
492        let (id1, id2): (Uuid, Uuid) = collector
493            .await
494            .expect("collector should complete");
495        assert_ne!(id1, id2, "Task IDs should be unique");
496    }
497
498    // -------------------------------------------------------------------------
499    // SolveTask Tests (internal type used by queue)
500    // -------------------------------------------------------------------------
501
502    #[test]
503    fn test_solve_task_wait_time_increases() {
504        let (response_tx, _response_rx) = oneshot::channel();
505        let task = SolveTask::new(Uuid::new_v4(), make_order(), response_tx, long_deadline());
506
507        let wait1 = task.wait_time();
508        std::thread::sleep(std::time::Duration::from_millis(10));
509        let wait2 = task.wait_time();
510
511        assert!(wait2 > wait1);
512    }
513
514    #[tokio::test]
515    async fn test_solve_task_respond_delivers_result() {
516        let (response_tx, response_rx) = oneshot::channel();
517        let task = SolveTask::new(Uuid::new_v4(), make_order(), response_tx, long_deadline());
518
519        task.respond(Ok(make_single_quote()));
520
521        let result = response_rx
522            .await
523            .expect("should receive response");
524        assert!(result.is_ok());
525    }
526
527    #[tokio::test]
528    async fn test_solve_task_respond_delivers_error() {
529        let (response_tx, response_rx) = oneshot::channel();
530        let task = SolveTask::new(Uuid::new_v4(), make_order(), response_tx, long_deadline());
531
532        task.respond(Err(SolveError::Timeout { elapsed_ms: 100 }));
533
534        let result = response_rx
535            .await
536            .expect("should receive response");
537        assert!(matches!(result, Err(SolveError::Timeout { elapsed_ms: 100 })));
538    }
539}