Initial Source code version
This commit is contained in:
parent
90f5f38812
commit
d979374cbc
23 changed files with 2628 additions and 0 deletions
95
app/workers/base_worker.py
Normal file
95
app/workers/base_worker.py
Normal file
|
|
@ -0,0 +1,95 @@
|
|||
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
|
||||
Loading…
Add table
Add a link
Reference in a new issue