use std::sync::Arc;
use tokio::sync::mpsc;
use tokio_util::sync::CancellationToken;
use rskit_errors::{AppError, AppResult, ErrorCode};
use rskit_provider::request_response_fn;
use rskit_provider::traits::RequestResponse;
use crate::bridge::{as_provider, from_provider};
use crate::event::{Event, Progress};
use crate::handler::Handler;
use crate::pool::{Pool, PoolConfig};
struct DoubleHandler;
#[async_trait::async_trait]
impl Handler<i32, i32> for DoubleHandler {
async fn handle(
&self,
task: i32,
_emit: mpsc::Sender<Event<i32>>,
_cancel: CancellationToken,
) -> AppResult<i32> {
Ok(task * 2)
}
}
struct ErrorHandler;
#[async_trait::async_trait]
impl Handler<i32, i32> for ErrorHandler {
async fn handle(
&self,
_task: i32,
_emit: mpsc::Sender<Event<i32>>,
_cancel: CancellationToken,
) -> AppResult<i32> {
Err(AppError::new(ErrorCode::Internal, "boom"))
}
}
struct ProgressHandler;
#[async_trait::async_trait]
impl Handler<i32, i32> for ProgressHandler {
async fn handle(
&self,
task: i32,
emit: mpsc::Sender<Event<i32>>,
_cancel: CancellationToken,
) -> AppResult<i32> {
let fake_id = uuid::Uuid::new_v4();
let p = Progress::new(1, Some(1));
let ev = Event::progress(fake_id, "progress-handler", p);
let _ = emit.send(ev).await;
Ok(task * 2)
}
}
struct PlusOneHandler;
#[async_trait::async_trait]
impl Handler<i32, i32> for PlusOneHandler {
async fn handle(
&self,
task: i32,
_emit: mpsc::Sender<Event<i32>>,
_cancel: CancellationToken,
) -> AppResult<i32> {
Ok(task + 1)
}
}
struct BlockingHandler {
started: Arc<tokio::sync::Notify>,
release: Arc<tokio::sync::Notify>,
}
#[async_trait::async_trait]
impl Handler<i32, i32> for BlockingHandler {
async fn handle(
&self,
task: i32,
_emit: mpsc::Sender<Event<i32>>,
_cancel: CancellationToken,
) -> AppResult<i32> {
self.started.notify_one();
self.release.notified().await;
Ok(task)
}
}
#[tokio::test]
async fn pool_submit_and_await_result() {
let pool = Pool::new(Arc::new(DoubleHandler), PoolConfig::new("test-pool"));
let handle = pool.submit(21).await.unwrap();
let result = handle.result().await;
assert_eq!(result.unwrap(), 42);
}
#[tokio::test]
async fn pool_submit_error_propagates() {
let pool = Pool::new(Arc::new(ErrorHandler), PoolConfig::new("error-pool"));
let handle = pool.submit(1).await.unwrap();
let result = handle.result().await;
assert!(result.is_err(), "expected Err but got {result:?}");
}
#[tokio::test]
async fn pool_task_handle_events() {
let pool = Pool::new(Arc::new(ProgressHandler), PoolConfig::new("event-pool"));
let handle = pool.submit(5).await.unwrap();
let task_id = handle.id;
let mut event_rx = handle.events();
let result = handle.result().await;
assert!(result.is_ok());
let mut received = Vec::new();
while let Ok(ev) = event_rx.try_recv() {
received.push(ev);
}
let has_result_event = received.iter().any(|ev| ev.task_id == task_id);
assert!(
has_result_event,
"expected at least one event with task_id {task_id}, got: {received:?}"
);
}
#[tokio::test]
async fn from_provider_bridge() {
let provider = Arc::new(request_response_fn("p", |x: i32| async move { Ok(x + 1) }));
let handler = from_provider(provider);
let pool = Pool::new(Arc::new(handler), PoolConfig::new("bridge-pool"));
let handle = pool.submit(9).await.unwrap();
let result = handle.result().await;
assert_eq!(result.unwrap(), 10);
}
#[tokio::test]
async fn as_provider_bridge() {
let handler: Arc<dyn Handler<i32, i32>> = Arc::new(PlusOneHandler);
let provider = as_provider("name", handler);
let result = provider.execute(5).await;
assert_eq!(result.unwrap(), 6);
}
#[tokio::test]
async fn reject_overflow_policy_rejects_new_submission() {
let started = Arc::new(tokio::sync::Notify::new());
let release = Arc::new(tokio::sync::Notify::new());
let pool = Pool::new(
Arc::new(BlockingHandler {
started: started.clone(),
release: release.clone(),
}),
PoolConfig::new("reject-pool")
.with_size(1)
.with_queue_size(1)
.with_overflow_policy(crate::OverflowPolicy::Reject),
);
let _first = pool.submit(1).await.unwrap();
started.notified().await;
let _second = pool.submit(2).await.unwrap();
let third = pool.submit(3).await;
assert!(third.is_err());
release.notify_waiters();
started.notified().await;
release.notify_waiters();
}
#[tokio::test]
async fn drop_oldest_overflow_policy_drops_queued_task() {
let started = Arc::new(tokio::sync::Notify::new());
let release = Arc::new(tokio::sync::Notify::new());
let pool = Pool::new(
Arc::new(BlockingHandler {
started: started.clone(),
release: release.clone(),
}),
PoolConfig::new("drop-oldest-pool")
.with_size(1)
.with_queue_size(1)
.with_overflow_policy(crate::OverflowPolicy::DropOldest),
);
let _first = pool.submit(1).await.unwrap();
started.notified().await;
let second = pool.submit(2).await.unwrap();
let _third = pool.submit(3).await.unwrap();
let dropped_error = second.result().await.unwrap_err();
assert_eq!(dropped_error.code(), ErrorCode::RateLimited);
release.notify_waiters();
started.notified().await;
release.notify_waiters();
}