use async_trait::async_trait;
use std::collections::HashMap;
use tokio::sync::RwLock;
use super::super::traits::TaskStorage;
use crate::error::Result;
use crate::task::{Task, TaskStatus};
pub struct InMemoryStorage {
tasks: RwLock<HashMap<String, Task>>,
}
impl InMemoryStorage {
pub fn new() -> Self {
Self {
tasks: RwLock::new(HashMap::new()),
}
}
}
#[async_trait]
impl TaskStorage for InMemoryStorage {
async fn save_task(&self, task: &Task) -> Result<()> {
let mut tasks = self.tasks.write().await;
tasks.insert(task.definition.id.clone(), task.clone());
Ok(())
}
async fn get_task(&self, task_id: &str) -> Result<Option<Task>> {
let tasks = self.tasks.read().await;
Ok(tasks.get(task_id).cloned())
}
async fn update_task_status(&self, task_id: &str, status: TaskStatus) -> Result<()> {
let mut tasks = self.tasks.write().await;
if let Some(task) = tasks.get_mut(task_id) {
task.status = status;
}
Ok(())
}
async fn list_tasks_by_status(&self, status: TaskStatus) -> Result<Vec<Task>> {
let tasks = self.tasks.read().await;
let filtered_tasks: Vec<Task> = tasks
.values()
.filter(|task| task.status == status)
.cloned()
.collect();
Ok(filtered_tasks)
}
async fn list_tasks_by_worker(&self, worker_id: &str) -> Result<Vec<Task>> {
let tasks = self.tasks.read().await;
let filtered_tasks: Vec<Task> = tasks
.values()
.filter(|task| task.assigned_worker.as_ref() == Some(&worker_id.to_string()))
.cloned()
.collect();
Ok(filtered_tasks)
}
async fn delete_task(&self, task_id: &str) -> Result<()> {
let mut tasks = self.tasks.write().await;
tasks.remove(task_id);
Ok(())
}
async fn get_pending_tasks(&self, limit: usize) -> Result<Vec<Task>> {
let mut tasks = self.list_tasks_by_status(TaskStatus::Pending).await?;
tasks.sort_by(|a, b| b.definition.priority.cmp(&a.definition.priority));
tasks.truncate(limit);
Ok(tasks)
}
async fn get_tasks_by_tags(&self, tags: &[String]) -> Result<Vec<Task>> {
let tasks = self.tasks.read().await;
let filtered_tasks: Vec<Task> = tasks
.values()
.filter(|task| task.definition.tags.iter().any(|tag| tags.contains(tag)))
.cloned()
.collect();
Ok(filtered_tasks)
}
}