Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
102 changes: 70 additions & 32 deletions ESSArch_Core/WorkflowEngine/dbtask.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
wait_random_exponential,
)

from ESSArch_Core.db.locks import RenewableCacheLock
from ESSArch_Core.db.utils import check_db_connection
from ESSArch_Core.essxml.Generator.xmlGenerator import parseContent
from ESSArch_Core.ip.models import EventIP, InformationPackage
Expand All @@ -47,33 +48,14 @@

User = get_user_model()

# import time
# from contextlib import contextmanager
#
# LOCK_EXPIRE = 60 * 10 # Lock expires in 10 minutes
#
# @contextmanager
# def cache_lock(lock_id):
# timeout_at = time.monotonic() + LOCK_EXPIRE - 3
# # cache.add fails if the key already exists
# # Second value is arbitrary
# status = cache.add(lock_id, "lock", timeout=LOCK_EXPIRE)
# try:
# yield status
# finally:
# if time.monotonic() < timeout_at and status:
# # don't release the lock if we exceeded the timeout
# # to lessen the chance of releasing an expired lock
# # owned by someone else
# # also don't release the lock if we didn't acquire it
# cache.delete(lock_id)


class DBTask(Task):
abstract = True
event_type = None
queue = 'celery'
track = True
lock_timeout = 300
lock_renew_interval = 60
logger = logging.getLogger('essarch')

def __call__(self, *args, **kwargs):
Expand Down Expand Up @@ -170,7 +152,12 @@ def _run(self, *args, **kwargs):
if self.parallel:
cm = nullcontext()
else:
cm = cache.lock(ip.get_lock_key(), timeout=300)
cm = RenewableCacheLock(
ip.get_lock_key(),
timeout=self.lock_timeout,
renew_interval=self.lock_renew_interval,
logger=self.logger,
)
try:
if ip.is_locked():
if not self.parallel:
Expand All @@ -183,7 +170,8 @@ def _run(self, *args, **kwargs):
ip, self.name, self.task_id))
with cm:
if not self.parallel:
self.logger.info('Task: {} ({}) acquired lock for IP {}'.format(self.name, self.task_id, ip))
self.logger.info('Task: {} ({}) acquired lock for IP {}'.format(
self.name, self.task_id, ip))
else:
self.logger.info('Task: {} ({}) is running in parallel for IP: {}'.format(
self.name, self.task_id, ip))
Expand All @@ -204,17 +192,32 @@ def _run(self, *args, **kwargs):
raise

if t.run_if and not self.parse_params(t.run_if)[0]:
self.logger.info(
'TASK SKIPPED: run_if=False task=%s id=%s',
self.name,
self.task_id,
)
r = None
t.hidden = True
t.save()
else:
r = self._run_task(*args, **kwargs)

if not self.parallel:
self.logger.info('{} released lock for IP: {}'.format(self.task_id, str(ip)))
self.logger.info(
'Task: %s (%s) released lock for IP %s',
self.name,
self.task_id,
ip,
)
except LockNotOwnedError:
self.logger.warning('Task: {} ({}) LockNotOwnedError for IP: {}'.format(
self.name, self.task_id, str(ip)))
r = None
self.logger.exception(
'Task: %s (%s) LockNotOwnedError for IP: %s',
self.name,
self.task_id,
str(ip),
)
raise
return r

return self._run_task(*args, **kwargs)
Expand All @@ -237,8 +240,13 @@ def _run_task(self, *args, **kwargs):
)
res = self.run(*args, **kwargs)
except exceptions.Ignore:
self.logger.warning("TASK IGNORED: %s (%s)", self.name, self.task_id)
raise
except exceptions.Retry as e:
self.logger.debug("TASK RETRY: %s (%s): %s", self.name, self.task_id, e)
raise
except Exception as e:
self.logger.warning("TASK EXCEPTION: %s (%s)", self.name, self.task_id)
einfo = ExceptionInfo()
self.failure(e, einfo)
if self.eager:
Expand All @@ -255,6 +263,7 @@ def _run_task(self, *args, **kwargs):

raise
else:
self.logger.info("TASK SUCCESS: %s (%s), result=%r", self.name, self.task_id, res)
self.success(res, args, kwargs)

return res
Expand All @@ -274,9 +283,15 @@ def after_return(self, status, retval, task_id, args, kwargs, einfo):
except ProcessStep.DoesNotExist:
return

# with cache_lock(step.cache_lock_key):
with cache.lock(step.cache_lock_key, timeout=60):
step.clear_cache()
try:
with cache.lock(step.cache_lock_key, timeout=60):
step.clear_cache()
except LockNotOwnedError:
self.logger.exception(
'Could not release step cache lock: task=%s step=%s',
self.task_id,
self.step,
)

return super().after_return(status, retval, task_id, args, kwargs, einfo)

Expand Down Expand Up @@ -316,6 +331,14 @@ def failure(self, exc, einfo):
timestamps
'''

self.logger.debug(
"CELERY FAILURE: task=%s id=%s eager=%s exception=%r",
self.name,
self.task_id,
self.eager,
exc,
)

if self.eager:
self.update_state(task_id=self.task_id, state=celery_states.FAILURE)
self.backend._store_result(
Expand All @@ -324,8 +347,15 @@ def failure(self, exc, einfo):
)

if self.event_type:
msg = einfo.traceback
self.create_event(celery_states.FAILURE, msg, None, einfo)
try:
msg = einfo.traceback
self.create_event(celery_states.FAILURE, msg, None, einfo)
except Exception:
self.logger.exception(
"Could not create failure event for task=%s id=%s",
self.name,
self.task_id,
)

def create_success_event(self, msg, retval=None):
return self.create_event(celery_states.SUCCESS, msg, retval, None)
Expand All @@ -347,6 +377,14 @@ def success(self, retval, args, kwargs):
of the current task but before the next task has started.
'''

self.logger.debug(
"CELERY SUCCESS: task=%s id=%s eager=%s result=%r",
self.name,
self.task_id,
self.eager,
retval,
)

if self.eager:
self.update_state(task_id=self.task_id, state=celery_states.SUCCESS)
self.backend.store_result(self.task_id, retval, celery_states.SUCCESS)
Expand Down
147 changes: 147 additions & 0 deletions ESSArch_Core/db/locks.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,147 @@
import logging
import threading

from django.core.cache import cache
from redis.exceptions import LockError, LockNotOwnedError

logger = logging.getLogger('essarch')


class RenewableCacheLock:
"""
A Redis lock with automatic TTL renewal.

The lock is acquired with a relatively short timeout and extended
periodically while the protected code is running.

The Redis lock uses thread_local=False because the lock is acquired
in the task thread and renewed from a separate thread.
"""

def __init__(
self,
key,
timeout=300,
renew_interval=None,
logger=None,
):
self.key = key
self.timeout = timeout

# Renew well before the TTL expires.
# Default to one third of the TTL, but never less than 1 second.
if renew_interval is None:
renew_interval = max(1, timeout // 3)

self.renew_interval = renew_interval
self.logger = logger or logging.getLogger('essarch')

self.lock = None
self._stop_event = threading.Event()
self._renew_thread = None
self._lost = False

def acquire(self):
self.lock = cache.lock(
self.key,
timeout=self.timeout,
thread_local=False,
)

acquired = self.lock.acquire()

if not acquired:
raise LockError(
'Could not acquire lock {}'.format(self.key)
)

self.logger.debug(
'Acquired lock %s with TTL %ss, renewing every %ss',
self.key,
self.timeout,
self.renew_interval,
)

self._renew_thread = threading.Thread(
target=self._renew_loop,
name='essarch-lock-renewer',
daemon=True,
)
self._renew_thread.start()

return True

def _renew_loop(self):
while not self._stop_event.wait(self.renew_interval):
try:
self.lock.extend(
self.timeout,
replace_ttl=True,
)

self.logger.debug(
'Renewed lock %s for %ss',
self.key,
self.timeout,
)

except LockNotOwnedError:
self._lost = True

self.logger.error(
'Lost ownership of lock %s while renewing',
self.key,
)
return

except Exception:
self._lost = True

self.logger.exception(
'Unexpected error renewing lock %s',
self.key,
)
return

def release(self):
self._stop_event.set()

if self._renew_thread is not None:
self._renew_thread.join(
timeout=self.renew_interval + 1,
)

if self.lock is None:
return

try:
self.lock.release()

self.logger.debug(
'Released lock %s',
self.key,
)

except LockNotOwnedError:
self.logger.warning(
'Lock %s was no longer owned when releasing',
self.key,
)

def __enter__(self):
self.acquire()
return self

@property
def lost(self):
return self._lost

def __exit__(self, exc_type, exc_value, traceback):
self.release()

if self._lost and exc_type is None:
raise LockNotOwnedError(
'Lock {} was lost while task was running'.format(self.key)
)

return False
12 changes: 6 additions & 6 deletions requirements/base.txt
Original file line number Diff line number Diff line change
@@ -1,11 +1,11 @@
asgiref==3.9.1
boto3==1.43.74
boto3==1.43.101
celery[tblib]==5.6.3
cffi==2.1.1
channels==4.3.2
channels-redis==4.3.0
chardet==7.6.0
click==8.4.2
click==8.5.0
cryptography==45.0.7
daphne==4.2.3
dj-rest-auth[with-social]==7.2.0
Expand All @@ -17,14 +17,14 @@ django-csp==3.7
django-environ==0.11.2
django-filter==26.1
django-groups-manager==1.3.0
django-guardian==3.3.3
django-guardian==3.5.0
django-languages-plus==2.1.1
django-mptt==0.18.0
django-nested-inline==0.4.6
django-picklefield==3.4.0
django-redis==7.0.0
django-relativity==0.2.6
djangorestframework==3.18.0
djangorestframework==3.18.1
django-json-widget==1.1.1
django-proxy==1.2.2
django-rest-knox==5.0.2
Expand All @@ -40,15 +40,15 @@ elasticsearch-dsl==7.4.1
gevent==24.11.1 ; platform_system=='Windows'
glob2==0.7
jsonfield==3.2.0
lxml==6.1.2
lxml==6.1.3
msoffcrypto-tool==5.4.2
natsort==8.4.0
opf-fido==1.6.1
pyfakefs==6.2.0
python-dateutil==2.8.2
pywin32==312 ; platform_system=='Windows'
redis==8.1.0
regex==2026.7.19
regex==2026.9.10
requests==2.34.2
requests-toolbelt==1.0.0
setuptools==81.0.0
Expand Down
Loading
Loading