1use std::time::Instant;
8
9use tokio::sync::oneshot;
10use uuid::Uuid;
11
12use crate::{types::internal::SolveTask, Order, SingleOrderQuote, SolveError, SolveParams};
13
14#[derive(Debug, Clone)]
16pub struct TaskQueueConfig {
17 pub capacity: usize,
19}
20
21impl Default for TaskQueueConfig {
22 fn default() -> Self {
23 Self { capacity: 1000 }
24 }
25}
26
27#[derive(Clone)]
31pub struct TaskQueueHandle {
32 sender: async_channel::Sender<SolveTask>,
33}
34
35impl TaskQueueHandle {
36 pub async fn enqueue(
40 &self,
41 order: Order,
42 params: SolveParams,
43 deadline: Instant,
44 ) -> Result<SingleOrderQuote, SolveError> {
45 let (response_tx, response_rx) = oneshot::channel();
47
48 let task_id = Uuid::new_v4();
50
51 let task = SolveTask::new(task_id, order, response_tx, deadline).with_params(params);
53
54 self.sender
56 .send(task)
57 .await
58 .map_err(|_| SolveError::QueueFull)?;
59
60 response_rx
62 .await
63 .map_err(|_| SolveError::Internal("worker dropped response channel".to_string()))?
64 }
65
66 #[cfg(test)]
70 pub fn approximate_depth(&self) -> usize {
71 self.sender.len()
72 }
73
74 #[cfg(test)]
76 pub fn is_full(&self) -> bool {
77 self.sender.is_full()
78 }
79
80 pub fn from_sender(sender: async_channel::Sender<SolveTask>) -> Self {
84 Self { sender }
85 }
86}
87
88pub struct TaskQueue {
92 receiver: async_channel::Receiver<SolveTask>,
93 handle: TaskQueueHandle,
94}
95
96impl TaskQueue {
97 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 pub fn split(self) -> (TaskQueueHandle, async_channel::Receiver<SolveTask>) {
107 (self.handle, self.receiver)
108 }
109
110 #[cfg(test)]
112 pub fn handle(&self) -> TaskQueueHandle {
113 self.handle.clone()
114 }
115
116 #[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 fn make_address(byte: u8) -> Address {
141 Address::from([byte; 20])
142 }
143
144 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 #[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 #[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 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 }
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 #[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 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 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 let worker = tokio::spawn(async move {
308 let task = receiver
309 .recv()
310 .await
311 .expect("should receive task");
312 drop(task); });
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);
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 #[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(); 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 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 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 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 let _task = receiver.recv().await.unwrap();
423
424 assert!(!handle.is_full());
426 assert_eq!(handle.approximate_depth(), 1);
427 }
428
429 #[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 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 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 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 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 #[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}