103 lines
3.4 KiB
Python
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)
|