1use parking_lot::RwLock;
4use serde::{Deserialize, Serialize};
5use std::collections::HashMap;
6use std::sync::Arc;
7use std::sync::atomic::{AtomicU64, Ordering};
8use tokio::sync::broadcast;
9
10type WorkId = String;
12
13#[derive(Debug, Clone, Serialize, Deserialize)]
15pub struct WorkItem {
16 pub id: WorkId,
18 pub work_type: String,
20 pub payload: serde_json::Value,
22 pub priority: i32,
24 pub status: WorkStatus,
26 pub claimed_by: Option<String>,
28 pub result: Option<WorkResult>,
30 pub created_at_ms: u64,
32 pub claimed_at_ms: Option<u64>,
34 pub completed_at_ms: Option<u64>,
36 pub max_retries: usize,
38 pub retry_count: usize,
40}
41
42#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
44pub enum WorkStatus {
45 Pending,
47 Claimed,
49 InProgress,
51 Completed,
53 Failed,
55 Cancelled,
57}
58
59#[derive(Debug, Clone, Serialize, Deserialize)]
61pub struct WorkResult {
62 pub success: bool,
64 pub content: String,
66 pub error: Option<String>,
68 pub duration_ms: u64,
70 pub tokens_used: Option<u64>,
72}
73
74#[derive(Debug, Clone, Serialize, Deserialize)]
76pub enum WorkEvent {
77 Enqueued {
79 id: WorkId,
81 work_type: String,
83 },
84 Claimed {
86 id: WorkId,
88 agent_id: String,
90 },
91 Started {
93 id: WorkId,
95 },
96 Completed {
98 id: WorkId,
100 success: bool,
102 },
103 Cancelled {
105 id: WorkId,
107 },
108}
109
110#[derive(Debug, Clone, Default, Serialize, Deserialize)]
112pub struct WorkQueueStats {
113 pub pending: usize,
115 pub claimed: usize,
117 pub in_progress: usize,
119 pub completed: usize,
121 pub failed: usize,
123 pub cancelled: usize,
125}
126
127#[derive(Debug, Clone)]
129pub struct WorkQueueConfig {
130 pub max_items: usize,
132}
133
134impl Default for WorkQueueConfig {
135 fn default() -> Self {
136 Self { max_items: 10_000 }
137 }
138}
139
140pub struct WorkQueue {
142 items: Arc<RwLock<HashMap<WorkId, WorkItem>>>,
143 next_id: AtomicU64,
144 #[allow(dead_code)]
145 config: WorkQueueConfig,
146 tx: broadcast::Sender<WorkEvent>,
147}
148
149impl WorkQueue {
150 pub fn new(config: WorkQueueConfig) -> Self {
152 let (tx, _) = broadcast::channel(256);
153 Self {
154 items: Arc::new(RwLock::new(HashMap::new())),
155 next_id: AtomicU64::new(1),
156 config,
157 tx,
158 }
159 }
160
161 pub fn enqueue(
163 &self,
164 work_type: impl Into<String>,
165 payload: serde_json::Value,
166 priority: i32,
167 ) -> WorkId {
168 let id = format!("wq-{}", self.next_id.fetch_add(1, Ordering::SeqCst));
169 let item = WorkItem {
170 id: id.clone(),
171 work_type: work_type.into(),
172 payload,
173 priority,
174 status: WorkStatus::Pending,
175 claimed_by: None,
176 result: None,
177 created_at_ms: now_ms(),
178 claimed_at_ms: None,
179 completed_at_ms: None,
180 max_retries: 3,
181 retry_count: 0,
182 };
183 self.items.write().insert(id.clone(), item);
184 let _ = self.tx.send(WorkEvent::Enqueued {
185 id: id.clone(),
186 work_type: String::new(),
187 });
188 id
189 }
190
191 pub fn claim(&self, agent_id: &str, work_type_filter: Option<&[String]>) -> Option<WorkItem> {
196 let mut items = self.items.write();
197 let mut best: Option<(WorkId, i32)> = None;
198
199 for (id, item) in items.iter() {
200 if item.status != WorkStatus::Pending {
201 continue;
202 }
203 if let Some(filter) = work_type_filter
204 && !filter.contains(&item.work_type)
205 {
206 continue;
207 }
208 match &best {
209 Some((_, best_pri)) if item.priority <= *best_pri => {}
210 _ => best = Some((id.clone(), item.priority)),
211 }
212 }
213
214 if let Some((id, _)) = best {
215 #[allow(clippy::unwrap_used)]
218 let item = items.get_mut(&id).unwrap();
219 item.status = WorkStatus::Claimed;
220 item.claimed_by = Some(agent_id.to_string());
221 item.claimed_at_ms = Some(now_ms());
222 let claimed = item.clone();
223 let _ = self.tx.send(WorkEvent::Claimed {
224 id: id.clone(),
225 agent_id: agent_id.to_string(),
226 });
227 Some(claimed)
228 } else {
229 None
230 }
231 }
232
233 pub fn start(&self, item_id: &str) -> crate::error::SdkResult<()> {
235 let mut items = self.items.write();
236 let item =
237 items
238 .get_mut(item_id)
239 .ok_or_else(|| crate::error::SdkError::WorkItemNotFound {
240 item_id: item_id.to_string(),
241 })?;
242 if item.status != WorkStatus::Claimed {
243 return Err(crate::error::SdkError::InvalidState {
244 entity: "work_item".into(),
245 reason: format!("item {} not in Claimed state", item_id),
246 });
247 }
248 item.status = WorkStatus::InProgress;
249 let _ = self.tx.send(WorkEvent::Started {
250 id: item_id.to_string(),
251 });
252 Ok(())
253 }
254
255 pub fn complete(&self, item_id: &str, result: WorkResult) -> crate::error::SdkResult<()> {
257 let mut items = self.items.write();
258 let item =
259 items
260 .get_mut(item_id)
261 .ok_or_else(|| crate::error::SdkError::WorkItemNotFound {
262 item_id: item_id.to_string(),
263 })?;
264 item.status = WorkStatus::Completed;
265 item.result = Some(result);
266 item.completed_at_ms = Some(now_ms());
267 let success = item.result.as_ref().map(|r| r.success).unwrap_or(false);
268 let _ = self.tx.send(WorkEvent::Completed {
269 id: item_id.to_string(),
270 success,
271 });
272 Ok(())
273 }
274
275 pub fn retry(&self, item_id: &str) -> anyhow::Result<bool> {
277 let mut items = self.items.write();
278 let item =
279 items
280 .get_mut(item_id)
281 .ok_or_else(|| crate::error::SdkError::WorkItemNotFound {
282 item_id: item_id.to_string(),
283 })?;
284 if item.retry_count >= item.max_retries {
285 return Ok(false);
286 }
287 item.retry_count += 1;
288 item.status = WorkStatus::Pending;
289 item.claimed_by = None;
290 item.claimed_at_ms = None;
291 Ok(true)
292 }
293
294 pub fn cancel(&self, item_id: &str) -> anyhow::Result<()> {
296 let mut items = self.items.write();
297 let item =
298 items
299 .get_mut(item_id)
300 .ok_or_else(|| crate::error::SdkError::WorkItemNotFound {
301 item_id: item_id.to_string(),
302 })?;
303 item.status = WorkStatus::Cancelled;
304 let _ = self.tx.send(WorkEvent::Cancelled {
305 id: item_id.to_string(),
306 });
307 Ok(())
308 }
309
310 pub fn get(&self, item_id: &str) -> Option<WorkItem> {
312 self.items.read().get(item_id).cloned()
313 }
314
315 pub fn list(&self, filter: Option<WorkStatus>) -> Vec<WorkItem> {
317 self.items
318 .read()
319 .values()
320 .filter(|item| match filter {
321 Some(s) => item.status == s,
322 None => true,
323 })
324 .cloned()
325 .collect()
326 }
327
328 pub fn stats(&self) -> WorkQueueStats {
330 let items = self.items.read();
331 let mut stats = WorkQueueStats::default();
332 for item in items.values() {
333 match item.status {
334 WorkStatus::Pending => stats.pending += 1,
335 WorkStatus::Claimed => stats.claimed += 1,
336 WorkStatus::InProgress => stats.in_progress += 1,
337 WorkStatus::Completed => stats.completed += 1,
338 WorkStatus::Failed => stats.failed += 1,
339 WorkStatus::Cancelled => stats.cancelled += 1,
340 }
341 }
342 stats
343 }
344
345 pub fn subscribe(&self) -> broadcast::Receiver<WorkEvent> {
347 self.tx.subscribe()
348 }
349}
350
351fn now_ms() -> u64 {
352 std::time::SystemTime::now()
353 .duration_since(std::time::UNIX_EPOCH)
354 .map(|d| d.as_millis() as u64)
355 .unwrap_or(0)
356}
357
358#[cfg(test)]
359mod tests {
360 use super::*;
361
362 #[test]
363 fn enqueue_and_claim() {
364 let q = WorkQueue::new(WorkQueueConfig::default());
365 let id = q.enqueue("review", serde_json::json!({"file": "main.rs"}), 1);
366 let item = q.claim("agent-1", None).unwrap();
367 assert_eq!(item.id, id);
368 assert_eq!(item.claimed_by.unwrap(), "agent-1");
369 assert_eq!(item.status, WorkStatus::Claimed);
370 }
371
372 #[test]
373 fn claim_is_atomic() {
374 let q = WorkQueue::new(WorkQueueConfig::default());
375 q.enqueue("task", serde_json::json!({}), 0);
376 let first = q.claim("a1", None);
377 let second = q.claim("a2", None);
378 assert!(first.is_some());
379 assert!(second.is_none());
380 }
381
382 #[test]
383 fn claim_respects_priority() {
384 let q = WorkQueue::new(WorkQueueConfig::default());
385 q.enqueue("low", serde_json::json!({}), 1);
386 q.enqueue("high", serde_json::json!({}), 10);
387 let item = q.claim("a1", None).unwrap();
388 assert_eq!(item.priority, 10);
389 }
390
391 #[test]
392 fn claim_with_type_filter() {
393 let q = WorkQueue::new(WorkQueueConfig::default());
394 q.enqueue("review", serde_json::json!({}), 0);
395 q.enqueue("build", serde_json::json!({}), 0);
396 let item = q.claim("a1", Some(&["build".into()]));
397 assert!(item.is_some());
398 assert_eq!(item.unwrap().work_type, "build");
399 }
400
401 #[test]
402 fn complete_item() {
403 let q = WorkQueue::new(WorkQueueConfig::default());
404 let id = q.enqueue("task", serde_json::json!({}), 0);
405 let _item = q.claim("a1", None).unwrap();
406 q.start(&id).unwrap();
407 q.complete(
408 &id,
409 WorkResult {
410 success: true,
411 content: "done".into(),
412 error: None,
413 duration_ms: 100,
414 tokens_used: None,
415 },
416 )
417 .unwrap();
418 let item = q.get(&id).unwrap();
419 assert_eq!(item.status, WorkStatus::Completed);
420 assert!(item.result.unwrap().success);
421 }
422
423 #[test]
424 fn retry_item() {
425 let q = WorkQueue::new(WorkQueueConfig::default());
426 let id = q.enqueue("task", serde_json::json!({}), 0);
427 {
428 let mut items = q.items.write();
429 let item = items.get_mut(&id).unwrap();
430 item.status = WorkStatus::Failed;
431 item.max_retries = 2;
432 }
433 assert!(q.retry(&id).unwrap());
434 let item = q.get(&id).unwrap();
435 assert_eq!(item.status, WorkStatus::Pending);
436 assert_eq!(item.retry_count, 1);
437 }
438
439 #[test]
440 fn cancel_item() {
441 let q = WorkQueue::new(WorkQueueConfig::default());
442 let id = q.enqueue("task", serde_json::json!({}), 0);
443 q.cancel(&id).unwrap();
444 assert_eq!(q.get(&id).unwrap().status, WorkStatus::Cancelled);
445 }
446
447 #[test]
448 fn queue_stats() {
449 let q = WorkQueue::new(WorkQueueConfig::default());
450 q.enqueue("t1", serde_json::json!({}), 0);
451 q.enqueue("t2", serde_json::json!({}), 0);
452 q.claim("a1", None);
453 let stats = q.stats();
454 assert_eq!(stats.pending, 1);
455 assert_eq!(stats.claimed, 1);
456 }
457
458 #[test]
459 fn subscribe_events() {
460 let q = WorkQueue::new(WorkQueueConfig::default());
461 let mut rx = q.subscribe();
462 q.enqueue("task", serde_json::json!({}), 0);
463 let event = rx.try_recv().unwrap();
464 assert!(matches!(event, WorkEvent::Enqueued { .. }));
465 }
466}