chore: add token revocation, caching, concurrency, and WebDAV lock modules
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
from .queue import TaskQueue, get_task_queue
|
||||
|
||||
__all__ = ["TaskQueue", "get_task_queue"]
|
||||
@@ -0,0 +1,251 @@
|
||||
import asyncio
|
||||
import time
|
||||
import uuid
|
||||
from typing import Dict, List, Any, Optional, Callable, Awaitable
|
||||
from dataclasses import dataclass, field
|
||||
from enum import Enum
|
||||
from collections import deque
|
||||
import logging
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class TaskStatus(Enum):
|
||||
PENDING = "pending"
|
||||
RUNNING = "running"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
|
||||
|
||||
class TaskPriority(Enum):
|
||||
LOW = 0
|
||||
NORMAL = 1
|
||||
HIGH = 2
|
||||
CRITICAL = 3
|
||||
|
||||
|
||||
@dataclass
|
||||
class Task:
|
||||
id: str
|
||||
queue_name: str
|
||||
handler_name: str
|
||||
payload: Dict[str, Any]
|
||||
priority: TaskPriority = TaskPriority.NORMAL
|
||||
status: TaskStatus = TaskStatus.PENDING
|
||||
created_at: float = field(default_factory=time.time)
|
||||
started_at: Optional[float] = None
|
||||
completed_at: Optional[float] = None
|
||||
result: Optional[Any] = None
|
||||
error: Optional[str] = None
|
||||
retry_count: int = 0
|
||||
max_retries: int = 3
|
||||
|
||||
|
||||
class TaskQueue:
|
||||
QUEUES = [
|
||||
"thumbnails",
|
||||
"cleanup",
|
||||
"billing",
|
||||
"notifications",
|
||||
"default",
|
||||
]
|
||||
|
||||
def __init__(self, max_workers: int = 4):
|
||||
self.max_workers = max_workers
|
||||
self.queues: Dict[str, deque] = {q: deque() for q in self.QUEUES}
|
||||
self.handlers: Dict[str, Callable] = {}
|
||||
self.tasks: Dict[str, Task] = {}
|
||||
self.workers: List[asyncio.Task] = []
|
||||
self._lock = asyncio.Lock()
|
||||
self._running = False
|
||||
self._stats = {
|
||||
"enqueued": 0,
|
||||
"completed": 0,
|
||||
"failed": 0,
|
||||
"retried": 0,
|
||||
}
|
||||
|
||||
async def start(self):
|
||||
self._running = True
|
||||
for i in range(self.max_workers):
|
||||
worker = asyncio.create_task(self._worker_loop(i))
|
||||
self.workers.append(worker)
|
||||
logger.info(f"TaskQueue started with {self.max_workers} workers")
|
||||
|
||||
async def stop(self):
|
||||
self._running = False
|
||||
for worker in self.workers:
|
||||
worker.cancel()
|
||||
await asyncio.gather(*self.workers, return_exceptions=True)
|
||||
self.workers.clear()
|
||||
logger.info("TaskQueue stopped")
|
||||
|
||||
def register_handler(self, name: str, handler: Callable[..., Awaitable[Any]]):
|
||||
self.handlers[name] = handler
|
||||
logger.debug(f"Registered handler: {name}")
|
||||
|
||||
async def enqueue(
|
||||
self,
|
||||
queue_name: str,
|
||||
handler_name: str,
|
||||
payload: Dict[str, Any],
|
||||
priority: TaskPriority = TaskPriority.NORMAL,
|
||||
max_retries: int = 3
|
||||
) -> str:
|
||||
if queue_name not in self.queues:
|
||||
queue_name = "default"
|
||||
|
||||
task_id = str(uuid.uuid4())
|
||||
task = Task(
|
||||
id=task_id,
|
||||
queue_name=queue_name,
|
||||
handler_name=handler_name,
|
||||
payload=payload,
|
||||
priority=priority,
|
||||
max_retries=max_retries
|
||||
)
|
||||
|
||||
async with self._lock:
|
||||
self.tasks[task_id] = task
|
||||
if priority == TaskPriority.HIGH or priority == TaskPriority.CRITICAL:
|
||||
self.queues[queue_name].appendleft(task_id)
|
||||
else:
|
||||
self.queues[queue_name].append(task_id)
|
||||
self._stats["enqueued"] += 1
|
||||
|
||||
logger.debug(f"Enqueued task {task_id} to {queue_name}")
|
||||
return task_id
|
||||
|
||||
async def get_task_status(self, task_id: str) -> Optional[Task]:
|
||||
async with self._lock:
|
||||
return self.tasks.get(task_id)
|
||||
|
||||
async def cancel_task(self, task_id: str) -> bool:
|
||||
async with self._lock:
|
||||
if task_id in self.tasks:
|
||||
task = self.tasks[task_id]
|
||||
if task.status == TaskStatus.PENDING:
|
||||
task.status = TaskStatus.CANCELLED
|
||||
return True
|
||||
return False
|
||||
|
||||
async def _worker_loop(self, worker_id: int):
|
||||
logger.debug(f"Worker {worker_id} started")
|
||||
while self._running:
|
||||
try:
|
||||
task = await self._get_next_task()
|
||||
if task:
|
||||
await self._process_task(task, worker_id)
|
||||
else:
|
||||
await asyncio.sleep(0.1)
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
logger.error(f"Worker {worker_id} error: {e}")
|
||||
await asyncio.sleep(1)
|
||||
logger.debug(f"Worker {worker_id} stopped")
|
||||
|
||||
async def _get_next_task(self) -> Optional[Task]:
|
||||
async with self._lock:
|
||||
for priority in [TaskPriority.CRITICAL, TaskPriority.HIGH, TaskPriority.NORMAL, TaskPriority.LOW]:
|
||||
for queue_name in self.QUEUES:
|
||||
if self.queues[queue_name]:
|
||||
task_id = None
|
||||
for tid in list(self.queues[queue_name]):
|
||||
if tid in self.tasks and self.tasks[tid].priority == priority:
|
||||
task_id = tid
|
||||
break
|
||||
if task_id:
|
||||
self.queues[queue_name].remove(task_id)
|
||||
task = self.tasks.get(task_id)
|
||||
if task and task.status == TaskStatus.PENDING:
|
||||
task.status = TaskStatus.RUNNING
|
||||
task.started_at = time.time()
|
||||
return task
|
||||
return None
|
||||
|
||||
async def _process_task(self, task: Task, worker_id: int):
|
||||
handler = self.handlers.get(task.handler_name)
|
||||
if not handler:
|
||||
logger.error(f"No handler found for {task.handler_name}")
|
||||
task.status = TaskStatus.FAILED
|
||||
task.error = f"Handler not found: {task.handler_name}"
|
||||
async with self._lock:
|
||||
self._stats["failed"] += 1
|
||||
return
|
||||
|
||||
try:
|
||||
result = await handler(**task.payload)
|
||||
task.status = TaskStatus.COMPLETED
|
||||
task.completed_at = time.time()
|
||||
task.result = result
|
||||
async with self._lock:
|
||||
self._stats["completed"] += 1
|
||||
logger.debug(f"Task {task.id} completed by worker {worker_id}")
|
||||
except Exception as e:
|
||||
task.error = str(e)
|
||||
task.retry_count += 1
|
||||
|
||||
if task.retry_count < task.max_retries:
|
||||
task.status = TaskStatus.PENDING
|
||||
task.started_at = None
|
||||
async with self._lock:
|
||||
self.queues[task.queue_name].append(task.id)
|
||||
self._stats["retried"] += 1
|
||||
logger.warning(f"Task {task.id} failed, retrying ({task.retry_count}/{task.max_retries})")
|
||||
else:
|
||||
task.status = TaskStatus.FAILED
|
||||
task.completed_at = time.time()
|
||||
async with self._lock:
|
||||
self._stats["failed"] += 1
|
||||
logger.error(f"Task {task.id} failed permanently: {e}")
|
||||
|
||||
async def cleanup_completed_tasks(self, max_age: float = 3600):
|
||||
now = time.time()
|
||||
async with self._lock:
|
||||
to_remove = [
|
||||
tid for tid, task in self.tasks.items()
|
||||
if task.status in [TaskStatus.COMPLETED, TaskStatus.FAILED, TaskStatus.CANCELLED]
|
||||
and task.completed_at and now - task.completed_at > max_age
|
||||
]
|
||||
for tid in to_remove:
|
||||
del self.tasks[tid]
|
||||
if to_remove:
|
||||
logger.debug(f"Cleaned up {len(to_remove)} completed tasks")
|
||||
|
||||
async def get_stats(self) -> Dict:
|
||||
async with self._lock:
|
||||
queue_sizes = {name: len(queue) for name, queue in self.queues.items()}
|
||||
pending = sum(1 for t in self.tasks.values() if t.status == TaskStatus.PENDING)
|
||||
running = sum(1 for t in self.tasks.values() if t.status == TaskStatus.RUNNING)
|
||||
return {
|
||||
**self._stats,
|
||||
"queue_sizes": queue_sizes,
|
||||
"pending_tasks": pending,
|
||||
"running_tasks": running,
|
||||
"total_tasks": len(self.tasks),
|
||||
}
|
||||
|
||||
|
||||
_task_queue: Optional[TaskQueue] = None
|
||||
|
||||
|
||||
async def init_task_queue(max_workers: int = 4) -> TaskQueue:
|
||||
global _task_queue
|
||||
_task_queue = TaskQueue(max_workers=max_workers)
|
||||
await _task_queue.start()
|
||||
return _task_queue
|
||||
|
||||
|
||||
async def shutdown_task_queue():
|
||||
global _task_queue
|
||||
if _task_queue:
|
||||
await _task_queue.stop()
|
||||
_task_queue = None
|
||||
|
||||
|
||||
def get_task_queue() -> TaskQueue:
|
||||
if not _task_queue:
|
||||
raise RuntimeError("Task queue not initialized")
|
||||
return _task_queue
|
||||
Reference in New Issue
Block a user