HR-ATS-Portal/backend/taskiq_management/middleware.py

103 lines
3.4 KiB
Python

"""PermanentTaskError + Redis Stream DLQ middleware for Taskiq.
Middleware order: DeadLetterMiddleware before SmartRetryMiddleware so
permanent failures can set retry_on_error=False before SmartRetry runs.
Pure module: no FastAPI imports.
"""
from __future__ import annotations
import json
import logging
from typing import Any
import redis.asyncio as redis
from taskiq import TaskiqMiddleware
from taskiq.message import TaskiqMessage
from taskiq.result import TaskiqResult
from taskiq_management.models import DLQ_STREAM
from taskiq_management.serializers import serialize_dlq_payload
logger=logging.getLogger("taskiq.dlq")
class PermanentTaskError(Exception):
"""Validation / business failure — DLQ immediately, no retries."""
class DeadLetterMiddleware(TaskiqMiddleware):
def __init__(self,redis_url:str,stream:str=DLQ_STREAM):
super().__init__()
self.redis_url=redis_url
self.stream=stream
self._redis:redis.Redis|None=None
async def startup(self) -> None:
self._redis=redis.from_url(self.redis_url,decode_responses=True)
async def shutdown(self) -> None:
if self._redis is not None:
await self._redis.aclose()
self._redis=None
def _client(self) -> redis.Redis:
if self._redis is None:
self._redis=redis.from_url(self.redis_url,decode_responses=True)
return self._redis
async def on_error(
self,
message:TaskiqMessage,
result:TaskiqResult[Any],
exception:BaseException,
) -> None:
retries=int(message.labels.get("_retries",0))
max_retries=int(message.labels.get("max_retries",2))
is_permanent=isinstance(exception,PermanentTaskError)
retries_exhausted=(retries+1)>=max_retries
if is_permanent:
message.labels["retry_on_error"]=False
if not is_permanent and not retries_exhausted:
return
queue=getattr(self.broker,"queue_name",None)
payload=serialize_dlq_payload(message,exception,retries=retries+1,queue=queue)
try:
await self._client().xadd(self.stream,{"payload":json.dumps(payload,ensure_ascii=False,default=str)})
logger.error(
"task %s (%s) sent to DLQ after %s",
message.task_name,
message.task_id,
"permanent failure" if is_permanent else f"{retries+1} attempts",
)
except Exception:
logger.exception("failed to write DLQ entry for %s",message.task_id)
await self._mark_inbox_dlq(message,exception)
async def _mark_inbox_dlq(self,message:TaskiqMessage,exception:BaseException) -> None:
if message.task_name!="inbox.match_message":
return
record_id=(message.kwargs or {}).get("record_id")
if not record_id and message.args:
record_id=message.args[0]
if not record_id:
return
try:
from db_setup import session_scope
from inbox.models import Inbox_Messages
async with session_scope() as session:
await Inbox_Messages.set_match_result(
session,
record_id,
status="dlq",
error=f"{type(exception).__name__}: {exception}",
)
except Exception:
logger.exception("failed to mark inbox %s as dlq",record_id)