import json import asyncio from uuid import UUID from datetime import datetime import aio_pika from sqlalchemy import update, select from app.db import SessionLocal from app.models import Task, TaskStatus from app.services import refresh_job_status, release_dependent_tasks from app.config import settings class BaseWorker: def __init__(self, routing_key: str, worker_id: str): self.routing_key = routing_key self.worker_id = worker_id async def run(self): connection = await aio_pika.connect_robust(settings.RABBITMQ_URL) async with connection: channel = await connection.channel() await channel.set_qos(prefetch_count=1) exchange = await channel.declare_exchange( settings.RABBITMQ_EXCHANGE, aio_pika.ExchangeType.TOPIC, durable=True ) queue = await channel.declare_queue( name=f"q.{self.routing_key}", durable=True, arguments={ "x-dead-letter-exchange": f"{settings.RABBITMQ_EXCHANGE}.dlx" }, ) await queue.bind(exchange, routing_key=self.routing_key) print(f"[{self.worker_id}] listening on {self.routing_key} ...") await queue.consume(self._on_message, no_ack=False) await asyncio.Future() async def _on_message(self, message: aio_pika.IncomingMessage): async with message.process(requeue=False): data = json.loads(message.body.decode("utf-8")) task_id = UUID(data["task_id"]) job_id = UUID(data["job_id"]) payload = data.get("payload", {}) try: await self._set_status(task_id, TaskStatus.RUNNING, started_at=datetime.utcnow()) result = await self.process(payload, job_id=job_id, task_id=task_id) await self._set_status(task_id, TaskStatus.SUCCESS, result=result, finished_at=datetime.utcnow()) async with SessionLocal() as session: await release_dependent_tasks(session, task_id) await refresh_job_status(session, job_id) await session.commit() print(f"[{self.worker_id}] SUCCESS {task_id}") except Exception as e: print(f"[{self.worker_id}] ERROR: {e}") await self._handle_failure(message, str(e), job_id) async def _handle_failure(self, message: aio_pika.IncomingMessage, err: str, job_id): data = json.loads(message.body.decode("utf-8")) task_id = UUID(data["task_id"]) async with SessionLocal() as session: res = await session.execute(select(Task).where(Task.id == task_id)) task = res.scalar_one_or_none() if task is None: return new_retries = (task.retries or 0) + 1 if new_retries <= (task.max_retries or settings.MAX_RETRIES): await session.execute( update(Task) .where(Task.id == task_id) .values(status=TaskStatus.QUEUED, retries=new_retries, error=f"Retry {new_retries}: {err}") ) await session.commit() await asyncio.sleep(min(2 ** new_retries, 30)) connection = await aio_pika.connect_robust(settings.RABBITMQ_URL) async with connection: ch = await connection.channel() ex = await ch.declare_exchange(settings.RABBITMQ_EXCHANGE, aio_pika.ExchangeType.TOPIC, durable=True) await ex.publish(aio_pika.Message(body=message.body), routing_key=message.routing_key) else: await session.execute( update(Task) .where(Task.id == task_id) .values(status=TaskStatus.FAILED, error=f"Max retries exceeded: {err}", finished_at=datetime.utcnow()) ) await refresh_job_status(session, job_id) await session.commit() async def _set_status(self, task_id: UUID, status: TaskStatus, **extra): async with SessionLocal() as session: await session.execute( update(Task).where(Task.id == task_id).values(status=status, worker_id=self.worker_id, **extra) ) await session.commit() async def process(self, payload: dict, **meta) -> dict: raise NotImplementedError