use crossbeam::queue::ArrayQueue;
#[derive(Debug, Clone)]
pub struct TaskRequest {
pub task_id: u64,
pub domain_id: u32,
pub payload: TaskPayload,
pub deadline_ms: u64,
pub channel_id: u32,
pub submitted_at: std::time::Instant,
}
impl TaskRequest {
#[inline]
pub fn is_expired(&self) -> bool {
if self.deadline_ms == 0 {
return false;
}
self.submitted_at.elapsed().as_millis() as u64 >= self.deadline_ms
}
}
#[derive(Debug, Clone)]
pub struct TaskResult {
pub task_id: u64,
pub domain_id: u32,
pub result: Result<TaskOutput, TaskError>,
}
#[derive(Debug, Clone)]
pub enum TaskPayload {
DbQuery {
sql_data: Vec<u8>,
},
FileRead {
path_data: Vec<u8>,
},
ExternalRpc {
url_data: Vec<u8>,
body_data: Vec<u8>,
},
ComplexCompute {
input_data: Vec<u8>,
},
Custom {
kind: u32,
data: Vec<u8>,
},
}
impl TaskPayload {
#[inline]
pub fn kind(&self) -> u32 {
match self {
TaskPayload::DbQuery { .. } => 1,
TaskPayload::FileRead { .. } => 2,
TaskPayload::ExternalRpc { .. } => 3,
TaskPayload::ComplexCompute { .. } => 4,
TaskPayload::Custom { kind, .. } => *kind,
}
}
#[inline]
pub fn is_write(&self) -> bool {
matches!(
self,
TaskPayload::DbQuery { .. } | TaskPayload::ExternalRpc { .. }
)
}
#[inline]
pub fn approximate_cost(&self) -> u32 {
match self {
TaskPayload::DbQuery { sql_data } => (sql_data.len() as u32).saturating_mul(10),
TaskPayload::FileRead { path_data } => (path_data.len() as u32).saturating_mul(20),
TaskPayload::ExternalRpc { url_data, .. } => (url_data.len() as u32).saturating_mul(50),
TaskPayload::ComplexCompute { input_data } => (input_data.len() as u32).saturating_mul(100),
TaskPayload::Custom { data, .. } => (data.len() as u32).saturating_mul(30),
}
}
}
#[derive(Debug, Clone)]
pub struct TaskOutput {
pub data: Vec<u8>,
pub status_code: u16,
pub elapsed_ms: u32,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum TaskError {
ChannelFull,
Timeout {
task_id: u64,
elapsed_ms: u32,
},
Cancelled {
task_id: u64,
},
ResourceExhausted,
InvalidPayload {
kind: u32,
},
Internal(String),
}
impl std::fmt::Display for TaskError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
TaskError::ChannelFull => write!(f, "task channel is full"),
TaskError::Timeout { task_id, elapsed_ms } => {
write!(f, "task {task_id} timed out after {elapsed_ms}ms")
}
TaskError::Cancelled { task_id } => write!(f, "task {task_id} was cancelled"),
TaskError::ResourceExhausted => write!(f, "blocking worker pool exhausted"),
TaskError::InvalidPayload { kind } => write!(f, "invalid payload kind: {kind}"),
TaskError::Internal(msg) => write!(f, "internal error: {msg}"),
}
}
}
impl std::error::Error for TaskError {}
#[derive(Debug)]
pub struct TaskOffloadChannel<const CAP: usize> {
request_queue: ArrayQueue<TaskRequest>,
result_queue: ArrayQueue<TaskResult>,
pending_count: std::sync::atomic::AtomicU64,
completed_count: std::sync::atomic::AtomicU64,
channel_id: u32,
}
impl<const CAP: usize> TaskOffloadChannel<CAP> {
pub fn new(channel_id: u32) -> Self {
Self {
request_queue: ArrayQueue::new(CAP),
result_queue: ArrayQueue::new(CAP),
pending_count: std::sync::atomic::AtomicU64::new(0),
completed_count: std::sync::atomic::AtomicU64::new(0),
channel_id,
}
}
#[inline]
pub fn channel_id(&self) -> u32 {
self.channel_id
}
#[inline]
pub const fn capacity(&self) -> usize {
CAP
}
#[inline]
pub fn pending_count(&self) -> u64 {
self.pending_count.load(std::sync::atomic::Ordering::Relaxed)
}
#[inline]
pub fn completed_count(&self) -> u64 {
self.completed_count.load(std::sync::atomic::Ordering::Relaxed)
}
#[inline]
pub fn is_empty(&self) -> bool {
self.request_queue.is_empty()
}
#[inline]
pub fn len(&self) -> usize {
self.request_queue.len()
}
#[inline]
pub fn submit(&self, req: TaskRequest) -> Result<(), TaskError> {
self.request_queue.push(req).map_err(|_| TaskError::ChannelFull)?;
self.pending_count.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Ok(())
}
#[inline]
pub fn submit_batch(&self, reqs: Vec<TaskRequest>) -> Result<usize, Vec<TaskRequest>> {
let mut iter = reqs.into_iter();
let mut count = 0usize;
let mut remaining: Vec<TaskRequest> = Vec::new();
let mut failed = false;
for req in iter.by_ref() {
match self.request_queue.push(req) {
Ok(()) => count += 1,
Err(req_back) => {
failed = true;
remaining.push(req_back);
break;
}
}
}
if failed {
remaining.extend(iter);
}
if count > 0 {
self.pending_count
.fetch_add(count as u64, std::sync::atomic::Ordering::Relaxed);
}
if failed {
Err(remaining)
} else {
Ok(count)
}
}
#[inline]
pub fn fetch(&self) -> Option<TaskRequest> {
loop {
let req = match self.request_queue.pop() {
Some(r) => r,
None => return None,
};
if !req.is_expired() {
return Some(req);
}
let elapsed = req.submitted_at.elapsed().as_millis() as u32;
let result = TaskResult {
task_id: req.task_id,
domain_id: req.domain_id,
result: Err(TaskError::Timeout {
task_id: req.task_id,
elapsed_ms: elapsed,
}),
};
if self.result_queue.push(result).is_ok() {
self.completed_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Self::dec_pending(&self.pending_count, 1);
}
}
}
#[inline]
pub fn fetch_batch(&self, max: usize) -> Vec<TaskRequest> {
let mut batch = Vec::with_capacity(max);
while batch.len() < max {
match self.request_queue.pop() {
Some(req) => {
if !req.is_expired() {
batch.push(req);
} else {
let elapsed = req.submitted_at.elapsed().as_millis() as u32;
let result = TaskResult {
task_id: req.task_id,
domain_id: req.domain_id,
result: Err(TaskError::Timeout {
task_id: req.task_id,
elapsed_ms: elapsed,
}),
};
if self.result_queue.push(result).is_ok() {
self.completed_count
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Self::dec_pending(&self.pending_count, 1);
}
}
}
None => break,
}
}
batch
}
#[inline]
pub fn complete(&self, result: TaskResult) -> Result<(), TaskError> {
self.result_queue.push(result).map_err(|_| TaskError::ChannelFull)?;
self.completed_count.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
Self::dec_pending(&self.pending_count, 1);
Ok(())
}
#[inline]
pub fn complete_batch(&self, results: Vec<TaskResult>) -> Result<usize, Vec<TaskResult>> {
let mut iter = results.into_iter();
let mut count = 0usize;
let mut remaining: Vec<TaskResult> = Vec::new();
let mut failed = false;
for result in iter.by_ref() {
match self.result_queue.push(result) {
Ok(()) => count += 1,
Err(result_back) => {
failed = true;
remaining.push(result_back);
break;
}
}
}
if failed {
remaining.extend(iter);
}
if count > 0 {
self.completed_count
.fetch_add(count as u64, std::sync::atomic::Ordering::Relaxed);
Self::dec_pending(&self.pending_count, count as u64);
}
if failed {
Err(remaining)
} else {
Ok(count)
}
}
#[inline]
fn dec_pending(counter: &std::sync::atomic::AtomicU64, n: u64) {
let mut cur = counter.load(std::sync::atomic::Ordering::Relaxed);
loop {
let new = cur.saturating_sub(n);
match counter.compare_exchange_weak(
cur,
new,
std::sync::atomic::Ordering::Relaxed,
std::sync::atomic::Ordering::Relaxed,
) {
Ok(_) => return,
Err(actual) => cur = actual,
}
}
}
#[inline]
pub fn poll_result(&self) -> Option<TaskResult> {
self.result_queue.pop()
}
#[inline]
pub fn poll_results(&self, max: usize) -> Vec<TaskResult> {
let mut results = Vec::with_capacity(max);
for _ in 0..max {
match self.result_queue.pop() {
Some(result) => results.push(result),
None => break,
}
}
results
}
#[inline]
pub fn build_request(
channel_id: u32,
task_id: u64,
domain_id: u32,
payload: TaskPayload,
deadline_ms: u64,
) -> TaskRequest {
TaskRequest {
task_id,
domain_id,
payload,
deadline_ms,
channel_id,
submitted_at: std::time::Instant::now(),
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn make_payload() -> TaskPayload {
TaskPayload::DbQuery {
sql_data: b"SELECT 1".to_vec(),
}
}
fn make_request(task_id: u64) -> TaskRequest {
TaskRequest {
task_id,
domain_id: 1,
payload: make_payload(),
deadline_ms: 1000,
channel_id: 1,
submitted_at: std::time::Instant::now(),
}
}
fn make_result(task_id: u64) -> TaskResult {
TaskResult {
task_id,
domain_id: 1,
result: Ok(TaskOutput {
data: b"OK".to_vec(),
status_code: 200,
elapsed_ms: 5,
}),
}
}
#[test]
fn test_channel_submit_fetch() {
let channel = TaskOffloadChannel::<256>::new(1);
let req = make_request(1);
channel.submit(req).unwrap();
assert_eq!(channel.pending_count(), 1);
assert!(!channel.is_empty());
let fetched = channel.fetch().unwrap();
assert_eq!(fetched.task_id, 1);
assert!(channel.is_empty());
}
#[test]
fn test_channel_complete_poll() {
let channel = TaskOffloadChannel::<256>::new(1);
channel.submit(make_request(1)).unwrap();
channel.fetch().unwrap();
channel.complete(make_result(1)).unwrap();
let result = channel.poll_result().unwrap();
assert_eq!(result.task_id, 1);
assert!(result.result.is_ok());
assert_eq!(channel.completed_count(), 1);
}
#[test]
fn test_channel_batch_operations() {
let channel = TaskOffloadChannel::<256>::new(1);
let reqs: Vec<TaskRequest> = (0..10).map(make_request).collect();
let submitted = channel.submit_batch(reqs).unwrap();
assert_eq!(submitted, 10);
assert_eq!(channel.pending_count(), 10);
let fetched = channel.fetch_batch(5);
assert_eq!(fetched.len(), 5);
assert_eq!(channel.pending_count(), 10);
let results: Vec<TaskResult> = fetched
.iter()
.map(|r| make_result(r.task_id))
.collect();
channel.complete_batch(results).unwrap();
assert_eq!(channel.completed_count(), 5);
assert_eq!(channel.pending_count(), 5);
let polled = channel.poll_results(10);
assert_eq!(polled.len(), 5);
}
#[test]
fn test_channel_capacity_limit() {
let channel = TaskOffloadChannel::<4>::new(1);
for i in 0..4 {
channel.submit(make_request(i)).unwrap();
}
let result = channel.submit(make_request(99));
assert!(result.is_err());
assert!(matches!(result.unwrap_err(), TaskError::ChannelFull));
}
#[test]
fn test_task_payload_methods() {
let db = TaskPayload::DbQuery {
sql_data: b"SELECT *".to_vec(),
};
assert_eq!(db.kind(), 1);
assert!(db.is_write());
let file = TaskPayload::FileRead {
path_data: b"/etc/passwd".to_vec(),
};
assert_eq!(file.kind(), 2);
assert!(!file.is_write());
let compute = TaskPayload::ComplexCompute {
input_data: vec![0; 100],
};
assert_eq!(compute.approximate_cost(), 10000);
}
#[test]
fn test_task_error_display() {
let err = TaskError::Timeout {
task_id: 42,
elapsed_ms: 5000,
};
assert!(format!("{err}").contains("42"));
assert!(format!("{err}").contains("5000"));
let err = TaskError::ChannelFull;
assert!(format!("{err}").contains("full"));
let err = TaskError::Cancelled { task_id: 7 };
assert!(format!("{err}").contains("7"));
}
#[test]
fn test_channel_stats() {
let channel = TaskOffloadChannel::<256>::new(42);
assert_eq!(channel.channel_id(), 42);
assert_eq!(channel.capacity(), 256);
assert_eq!(channel.pending_count(), 0);
assert_eq!(channel.completed_count(), 0);
channel.submit(make_request(1)).unwrap();
assert_eq!(channel.pending_count(), 1);
}
#[test]
fn test_channel_empty_operations() {
let channel = TaskOffloadChannel::<16>::new(1);
assert!(channel.is_empty());
assert_eq!(channel.len(), 0);
assert!(channel.fetch().is_none());
assert!(channel.poll_result().is_none());
let batch = channel.fetch_batch(5);
assert!(batch.is_empty());
let results = channel.poll_results(5);
assert!(results.is_empty());
}
#[test]
fn test_build_request() {
let req = TaskOffloadChannel::<16>::build_request(
7,
100,
2,
TaskPayload::FileRead {
path_data: b"/tmp/test".to_vec(),
},
5000,
);
assert_eq!(req.task_id, 100);
assert_eq!(req.domain_id, 2);
assert_eq!(req.deadline_ms, 5000);
assert_eq!(req.channel_id, 7);
assert!(req.submitted_at.elapsed().as_secs() < 5);
}
#[test]
fn test_complete_without_submit_no_underflow() {
let channel = TaskOffloadChannel::<16>::new(1);
assert_eq!(channel.pending_count(), 0);
channel.complete(make_result(1)).unwrap();
assert_eq!(channel.pending_count(), 0);
assert_eq!(channel.completed_count(), 1);
channel.complete_batch(vec![make_result(2), make_result(3)]).unwrap();
assert_eq!(channel.pending_count(), 0);
assert_eq!(channel.completed_count(), 3);
}
#[test]
fn test_fetch_expired_task_produces_timeout() {
let channel = TaskOffloadChannel::<16>::new(1);
let mut req = make_request(1);
req.deadline_ms = 1;
req.submitted_at = std::time::Instant::now() - std::time::Duration::from_secs(60);
assert_eq!(req.is_expired(), true);
channel.submit(req).unwrap();
assert!(channel.fetch().is_none());
let result = channel.poll_result().unwrap();
assert!(matches!(result.result, Err(TaskError::Timeout { .. })));
assert_eq!(channel.pending_count(), 0);
assert_eq!(channel.completed_count(), 1);
}
#[test]
fn test_fetch_batch_skips_expired_keeps_valid() {
let channel = TaskOffloadChannel::<16>::new(1);
let mut expired = make_request(1);
expired.deadline_ms = 1;
expired.submitted_at = std::time::Instant::now() - std::time::Duration::from_secs(60);
channel.submit(expired).unwrap();
channel.submit(make_request(2)).unwrap();
let batch = channel.fetch_batch(8);
assert_eq!(batch.len(), 1);
assert_eq!(batch[0].task_id, 2);
let result = channel.poll_result().unwrap();
assert!(matches!(result.result, Err(TaskError::Timeout { .. })));
assert_eq!(channel.pending_count(), 1);
assert_eq!(channel.completed_count(), 1);
}
#[test]
fn test_expired_result_push_failure_does_not_unbalance_pending() {
let channel = TaskOffloadChannel::<4>::new(1);
for i in 0..4 {
channel.complete(make_result(i)).unwrap();
}
assert_eq!(channel.pending_count(), 0);
let mut expired = make_request(99);
expired.deadline_ms = 1;
expired.submitted_at = std::time::Instant::now() - std::time::Duration::from_secs(60);
channel.submit(expired).unwrap();
assert_eq!(channel.pending_count(), 1);
assert!(channel.fetch().is_none());
assert_eq!(
channel.pending_count(),
1,
"结果队列满导致 Timeout 结果丢失时,pending 不得回退(账本失衡回归)"
);
assert_eq!(
channel.completed_count(),
4,
"结果未落地不得计入 completed"
);
}
}