110 lines
4.0 KiB
Python
110 lines
4.0 KiB
Python
"""In-memory task state management."""
|
|
|
|
from typing import Dict, Optional, List
|
|
from datetime import datetime, timedelta
|
|
import threading
|
|
import time
|
|
|
|
from backend.config import settings
|
|
from backend.models.response import TaskStatus, TrackInfo
|
|
from backend.services.progress_publisher import publish_progress
|
|
|
|
|
|
class TaskManager:
|
|
def __init__(self):
|
|
self._tasks: Dict[str, dict] = {}
|
|
self._lock = threading.Lock()
|
|
|
|
def create_task(self, task_id: str, filename: str, file_size: int) -> None:
|
|
with self._lock:
|
|
self._tasks[task_id] = {
|
|
"task_id": task_id,
|
|
"filename": filename,
|
|
"file_size": file_size,
|
|
"status": TaskStatus.PENDING,
|
|
"progress": 0,
|
|
"message": "Upload complete",
|
|
"error": None,
|
|
"tracks": [],
|
|
"created_at": datetime.now(),
|
|
"updated_at": datetime.now()
|
|
}
|
|
|
|
def has_task(self, task_id: str) -> bool:
|
|
return task_id in self._tasks
|
|
|
|
def get_status(self, task_id: str) -> dict:
|
|
"""Get task status with datetime objects converted to ISO format."""
|
|
with self._lock:
|
|
if task_id not in self._tasks:
|
|
return {}
|
|
# Return a copy with datetime objects converted to ISO strings
|
|
status = self._tasks[task_id].copy()
|
|
return self._prepare_status_for_serialization(status)
|
|
|
|
def _prepare_status_for_serialization(self, status: dict) -> dict:
|
|
"""Convert datetime objects to ISO format strings for JSON serialization."""
|
|
serializable = {}
|
|
for key, value in status.items():
|
|
if isinstance(value, datetime):
|
|
serializable[key] = value.isoformat()
|
|
elif isinstance(value, list):
|
|
# Handle lists of objects (e.g., tracks)
|
|
serializable[key] = [
|
|
{k: v.isoformat() if isinstance(v, datetime) else v for k, v in item.items()}
|
|
if isinstance(item, dict)
|
|
else item
|
|
for item in value
|
|
]
|
|
elif isinstance(value, dict):
|
|
# Recursively handle nested dicts
|
|
serializable[key] = {
|
|
k: v.isoformat() if isinstance(v, datetime) else v
|
|
for k, v in value.items()
|
|
}
|
|
else:
|
|
serializable[key] = value
|
|
return serializable
|
|
|
|
def update_task(self, task_id: str, **kwargs) -> None:
|
|
with self._lock:
|
|
if task_id in self._tasks:
|
|
self._tasks[task_id].update(kwargs)
|
|
self._tasks[task_id]["updated_at"] = datetime.now()
|
|
|
|
def update_task_with_progress(self, task_id: str, progress: int, message: str,
|
|
status: str = None, error: str = None) -> None:
|
|
update_kwargs = {
|
|
"progress": progress,
|
|
"message": message
|
|
}
|
|
if status:
|
|
update_kwargs["status"] = status
|
|
if error is not None:
|
|
update_kwargs["error"] = error
|
|
|
|
self.update_task(task_id, **update_kwargs)
|
|
|
|
# Publish progress via WebSocket
|
|
publish_progress(task_id, progress, message, status or "processing")
|
|
|
|
def cleanup_old_tasks(self) -> None:
|
|
now = datetime.now()
|
|
timeout = timedelta(seconds=settings.cleanup_after_seconds)
|
|
with self._lock:
|
|
to_delete = []
|
|
for task_id, data in self._tasks.items():
|
|
if now - data["created_at"] > timeout:
|
|
to_delete.append(task_id)
|
|
for task_id in to_delete:
|
|
self._delete_task_files(task_id)
|
|
del self._tasks[task_id]
|
|
|
|
def _delete_task_files(self, task_id: str) -> None:
|
|
import shutil
|
|
task_dir = settings.temp_dir / task_id
|
|
if task_dir.exists():
|
|
shutil.rmtree(task_dir, ignore_errors=True)
|
|
|
|
|
|
task_manager = TaskManager() |