Linting and hints

This commit is contained in:
Ilan Steemers
2020-06-12 12:33:04 +02:00
parent 2f2269c27f
commit ef4c7876ae
11 changed files with 148 additions and 125 deletions
+21 -18
View File
@@ -1,34 +1,37 @@
from django_q.conf import Conf
from django_q.brokers import Broker
from boto3 import Session from boto3 import Session
from django_q.brokers import Broker
from django_q.conf import Conf
class Sqs(Broker): class Sqs(Broker):
def __init__(self, list_key=Conf.PREFIX): def __init__(self, list_key: str = Conf.PREFIX):
self.sqs = None self.sqs = None
super(Sqs, self).__init__(list_key) super(Sqs, self).__init__(list_key)
self.queue = self.get_queue() self.queue = self.get_queue()
def enqueue(self, task): def enqueue(self, task):
response = self.queue.send_message(MessageBody=task) response = self.queue.send_message(MessageBody=task)
return response.get('MessageId') return response.get("MessageId")
def dequeue(self): def dequeue(self):
# sqs supports max 10 messages in bulk # sqs supports max 10 messages in bulk
if Conf.BULK > 10: if Conf.BULK > 10:
Conf.BULK = 10 Conf.BULK = 10
tasks = self.queue.receive_messages(MaxNumberOfMessages=Conf.BULK, VisibilityTimeout=Conf.RETRY) tasks = self.queue.receive_messages(
MaxNumberOfMessages=Conf.BULK, VisibilityTimeout=Conf.RETRY
)
if tasks: if tasks:
return [(t.receipt_handle, t.body) for t in tasks] return [(t.receipt_handle, t.body) for t in tasks]
def acknowledge(self, task_id): def acknowledge(self, task_id):
return self.delete(task_id) return self.delete(task_id)
def queue_size(self): def queue_size(self) -> int:
return int(self.queue.attributes['ApproximateNumberOfMessages']) return int(self.queue.attributes["ApproximateNumberOfMessages"])
def lock_size(self): def lock_size(self) -> int:
return int(self.queue.attributes['ApproximateNumberOfMessagesNotVisible']) return int(self.queue.attributes["ApproximateNumberOfMessagesNotVisible"])
def delete(self, task_id): def delete(self, task_id):
message = self.sqs.Message(self.queue.url, task_id) message = self.sqs.Message(self.queue.url, task_id)
@@ -43,20 +46,20 @@ class Sqs(Broker):
def purge_queue(self): def purge_queue(self):
self.queue.purge() self.queue.purge()
def ping(self): def ping(self) -> bool:
return 'sqs' in self.connection.get_available_resources() return "sqs" in self.connection.get_available_resources()
def info(self): def info(self) -> str:
return 'AWS SQS' return "AWS SQS"
@staticmethod @staticmethod
def get_connection(list_key=Conf.PREFIX): def get_connection(list_key: str = Conf.PREFIX) -> Session:
config = Conf.SQS config = Conf.SQS
if 'aws_region' in config: if "aws_region" in config:
config['region_name'] = config['aws_region'] config["region_name"] = config["aws_region"]
del(config['aws_region']) del config["aws_region"]
return Session(**config) return Session(**config)
def get_queue(self): def get_queue(self):
self.sqs = self.connection.resource('sqs') self.sqs = self.connection.resource("sqs")
return self.sqs.create_queue(QueueName=self.list_key) return self.sqs.create_queue(QueueName=self.list_key)
+7 -4
View File
@@ -1,5 +1,8 @@
import random import random
import redis import redis
from redis import Redis
from django_q.brokers import Broker from django_q.brokers import Broker
from django_q.conf import Conf from django_q.conf import Conf
@@ -25,7 +28,7 @@ class Disque(Broker):
command = "FASTACK" if Conf.DISQUE_FASTACK else "ACKJOB" command = "FASTACK" if Conf.DISQUE_FASTACK else "ACKJOB"
return self.connection.execute_command(f"{command} {task_id}") return self.connection.execute_command(f"{command} {task_id}")
def ping(self): def ping(self) -> bool:
return self.connection.execute_command("HELLO")[0] > 0 return self.connection.execute_command("HELLO")[0] > 0
def delete(self, task_id): def delete(self, task_id):
@@ -34,21 +37,21 @@ class Disque(Broker):
def fail(self, task_id): def fail(self, task_id):
return self.delete(task_id) return self.delete(task_id)
def delete_queue(self): def delete_queue(self) -> int:
jobs = self.connection.execute_command(f"JSCAN QUEUE {self.list_key}")[1] jobs = self.connection.execute_command(f"JSCAN QUEUE {self.list_key}")[1]
if jobs: if jobs:
job_ids = " ".join(jid.decode() for jid in jobs) job_ids = " ".join(jid.decode() for jid in jobs)
self.connection.execute_command(f"DELJOB {job_ids}") self.connection.execute_command(f"DELJOB {job_ids}")
return len(jobs) return len(jobs)
def info(self): def info(self) -> str:
if not self._info: if not self._info:
info = self.connection.info("server") info = self.connection.info("server")
self._info = f'Disque {info["disque_version"]}' self._info = f'Disque {info["disque_version"]}'
return self._info return self._info
@staticmethod @staticmethod
def get_connection(list_key=Conf.PREFIX): def get_connection(list_key: str = Conf.PREFIX) -> Redis:
# randomize nodes # randomize nodes
random.shuffle(Conf.DISQUE_NODES) random.shuffle(Conf.DISQUE_NODES)
# find one that works # find one that works
+12 -11
View File
@@ -1,31 +1,32 @@
from iron_mq import IronMQ, Queue
from requests.exceptions import HTTPError from requests.exceptions import HTTPError
from django_q.conf import Conf
from django_q.brokers import Broker from django_q.brokers import Broker
from iron_mq import IronMQ from django_q.conf import Conf
class IronMQBroker(Broker): class IronMQBroker(Broker):
def enqueue(self, task): def enqueue(self, task):
return self.connection.post(task)['ids'][0] return self.connection.post(task)["ids"][0]
def dequeue(self): def dequeue(self):
timeout = Conf.RETRY or None timeout = Conf.RETRY or None
tasks = self.connection.get(timeout=timeout, wait=1, max=Conf.BULK)['messages'] tasks = self.connection.get(timeout=timeout, wait=1, max=Conf.BULK)["messages"]
if tasks: if tasks:
return [(t['id'], t['body']) for t in tasks] return [(t["id"], t["body"]) for t in tasks]
def ping(self): def ping(self) -> bool:
return self.connection.name == self.list_key return self.connection.name == self.list_key
def info(self): def info(self) -> str:
return 'IronMQ' return "IronMQ"
def queue_size(self): def queue_size(self):
return self.connection.size() return self.connection.size()
def delete_queue(self): def delete_queue(self):
try: try:
return self.connection.delete_queue()['msg'] return self.connection.delete_queue()["msg"]
except HTTPError: except HTTPError:
return False return False
@@ -34,7 +35,7 @@ class IronMQBroker(Broker):
def delete(self, task_id): def delete(self, task_id):
try: try:
return self.connection.delete(task_id)['msg'] return self.connection.delete(task_id)["msg"]
except HTTPError: except HTTPError:
return False return False
@@ -45,6 +46,6 @@ class IronMQBroker(Broker):
return self.delete(task_id) return self.delete(task_id)
@staticmethod @staticmethod
def get_connection(list_key=Conf.PREFIX): def get_connection(list_key: str = Conf.PREFIX) -> Queue:
ironmq = IronMQ(name=None, **Conf.IRON_MQ) ironmq = IronMQ(name=None, **Conf.IRON_MQ)
return ironmq.queue(queue_name=list_key) return ironmq.queue(queue_name=list_key)
+3 -4
View File
@@ -4,7 +4,6 @@ from time import sleep
from bson import ObjectId from bson import ObjectId
from django.utils import timezone from django.utils import timezone
from pymongo import MongoClient from pymongo import MongoClient
from pymongo.errors import ConfigurationError from pymongo.errors import ConfigurationError
from django_q.brokers import Broker from django_q.brokers import Broker
@@ -21,7 +20,7 @@ class Mongo(Broker):
self.collection = self.get_collection() self.collection = self.get_collection()
@staticmethod @staticmethod
def get_connection(list_key=Conf.PREFIX): def get_connection(list_key: str = Conf.PREFIX) -> MongoClient:
return MongoClient(**Conf.MONGO) return MongoClient(**Conf.MONGO)
def get_collection(self): def get_collection(self):
@@ -41,10 +40,10 @@ class Mongo(Broker):
def purge_queue(self): def purge_queue(self):
return self.delete_queue() return self.delete_queue()
def ping(self): def ping(self) -> bool:
return self.info is not None return self.info is not None
def info(self): def info(self) -> str:
if not self._info: if not self._info:
self._info = f"MongoDB {self.connection.server_info()['version']}" self._info = f"MongoDB {self.connection.server_info()['version']}"
return self._info return self._info
+11 -9
View File
@@ -1,13 +1,13 @@
from datetime import timedelta from datetime import timedelta
from time import sleep from time import sleep
from django.utils import timezone
from django import db from django import db
from django.db import transaction from django.db import transaction
from django.utils import timezone
from django_q.brokers import Broker from django_q.brokers import Broker
from django_q.models import OrmQ
from django_q.conf import Conf, logger from django_q.conf import Conf, logger
from django_q.models import OrmQ
def _timeout(): def _timeout():
@@ -16,8 +16,10 @@ def _timeout():
class ORM(Broker): class ORM(Broker):
@staticmethod @staticmethod
def get_connection(list_key=Conf.PREFIX): def get_connection(list_key: str = Conf.PREFIX):
if transaction.get_autocommit(using=Conf.ORM): # Only True when not in an atomic block if transaction.get_autocommit(
using=Conf.ORM
): # Only True when not in an atomic block
# Make sure stale connections in the broker thread are explicitly # Make sure stale connections in the broker thread are explicitly
# closed before attempting DB access. # closed before attempting DB access.
# logger.debug("Broker thread calling close_old_connections") # logger.debug("Broker thread calling close_old_connections")
@@ -26,14 +28,14 @@ class ORM(Broker):
logger.debug("Broker in an atomic transaction") logger.debug("Broker in an atomic transaction")
return OrmQ.objects.using(Conf.ORM) return OrmQ.objects.using(Conf.ORM)
def queue_size(self): def queue_size(self) -> int:
return ( return (
self.get_connection() self.get_connection()
.filter(key=self.list_key, lock__lte=_timeout()) .filter(key=self.list_key, lock__lte=_timeout())
.count() .count()
) )
def lock_size(self): def lock_size(self) -> int:
return ( return (
self.get_connection().filter(key=self.list_key, lock__gt=_timeout()).count() self.get_connection().filter(key=self.list_key, lock__gt=_timeout()).count()
) )
@@ -41,10 +43,10 @@ class ORM(Broker):
def purge_queue(self): def purge_queue(self):
return self.get_connection().filter(key=self.list_key).delete() return self.get_connection().filter(key=self.list_key).delete()
def ping(self): def ping(self) -> bool:
return True return True
def info(self): def info(self) -> str:
if not self._info: if not self._info:
self._info = f"ORM {Conf.ORM}" self._info = f"ORM {Conf.ORM}"
return self._info return self._info
@@ -60,7 +62,7 @@ class ORM(Broker):
def dequeue(self): def dequeue(self):
tasks = self.get_connection().filter(key=self.list_key, lock__lt=_timeout())[ tasks = self.get_connection().filter(key=self.list_key, lock__lt=_timeout())[
0: Conf.BULK 0 : Conf.BULK
] ]
if tasks: if tasks:
task_list = [] task_list = []
+8 -7
View File
@@ -1,4 +1,5 @@
import redis import redis
from redis import Redis
from django_q.brokers import Broker from django_q.brokers import Broker
from django_q.conf import Conf, logger from django_q.conf import Conf, logger
@@ -10,7 +11,7 @@ except ImportError:
class Redis(Broker): class Redis(Broker):
def __init__(self, list_key=Conf.PREFIX): def __init__(self, list_key: str = Conf.PREFIX):
super(Redis, self).__init__(list_key=f"django_q:{list_key}:q") super(Redis, self).__init__(list_key=f"django_q:{list_key}:q")
def enqueue(self, task): def enqueue(self, task):
@@ -30,33 +31,33 @@ class Redis(Broker):
def purge_queue(self): def purge_queue(self):
return self.connection.ltrim(self.list_key, 1, 0) return self.connection.ltrim(self.list_key, 1, 0)
def ping(self): def ping(self) -> bool:
try: try:
return self.connection.ping() return self.connection.ping()
except redis.ConnectionError as e: except redis.ConnectionError as e:
logger.error("Can not connect to Redis server.") logger.error("Can not connect to Redis server.")
raise e raise e
def info(self): def info(self) -> str:
if not self._info: if not self._info:
info = self.connection.info("server") info = self.connection.info("server")
self._info = f"Redis {info['redis_version']}" self._info = f"Redis {info['redis_version']}"
return self._info return self._info
def set_stat(self, key, value, timeout): def set_stat(self, key: str, value: str, timeout: int):
self.connection.set(key, value, timeout) self.connection.set(key, value, timeout)
def get_stat(self, key): def get_stat(self, key: str):
if self.connection.exists(key): if self.connection.exists(key):
return self.connection.get(key) return self.connection.get(key)
def get_stats(self, pattern): def get_stats(self, pattern: str):
keys = self.connection.keys(pattern=pattern) keys = self.connection.keys(pattern=pattern)
if keys: if keys:
return self.connection.mget(keys) return self.connection.mget(keys)
@staticmethod @staticmethod
def get_connection(list_key=Conf.PREFIX): def get_connection(list_key: str = Conf.PREFIX) -> Redis:
if django_redis and Conf.DJANGO_REDIS: if django_redis and Conf.DJANGO_REDIS:
return django_redis.get_redis_connection(Conf.DJANGO_REDIS) return django_redis.get_redis_connection(Conf.DJANGO_REDIS)
if isinstance(Conf.REDIS, str): if isinstance(Conf.REDIS, str):
+36 -33
View File
@@ -1,5 +1,4 @@
import ast import ast
# Standard # Standard
import importlib import importlib
import signal import signal
@@ -11,7 +10,6 @@ from time import sleep
# external # external
import arrow import arrow
# Django # Django
from django import db from django import db
from django.conf import settings from django.conf import settings
@@ -20,7 +18,7 @@ from django.utils.translation import gettext_lazy as _
# Local # Local
import django_q.tasks import django_q.tasks
from django_q.brokers import get_broker from django_q.brokers import get_broker, Broker
from django_q.conf import Conf, logger, psutil, get_ppid, error_reporter from django_q.conf import Conf, logger, psutil, get_ppid, error_reporter
from django_q.humanhash import humanize from django_q.humanhash import humanize
from django_q.models import Task, Success, Schedule from django_q.models import Task, Success, Schedule
@@ -31,7 +29,7 @@ from django_q.status import Stat, Status
class Cluster: class Cluster:
def __init__(self, broker=None): def __init__(self, broker: Broker = None):
self.broker = broker or get_broker() self.broker = broker or get_broker()
self.sentinel = None self.sentinel = None
self.stop_event = None self.stop_event = None
@@ -43,7 +41,7 @@ class Cluster:
signal.signal(signal.SIGTERM, self.sig_handler) signal.signal(signal.SIGTERM, self.sig_handler)
signal.signal(signal.SIGINT, self.sig_handler) signal.signal(signal.SIGINT, self.sig_handler)
def start(self): def start(self) -> int:
# Start Sentinel # Start Sentinel
self.stop_event = Event() self.stop_event = Event()
self.start_event = Event() self.start_event = Event()
@@ -63,7 +61,7 @@ class Cluster:
sleep(0.1) sleep(0.1)
return self.pid return self.pid
def stop(self): def stop(self) -> bool:
if not self.sentinel.is_alive(): if not self.sentinel.is_alive():
return False return False
logger.info(_(f"Q Cluster {self.name} stopping.")) logger.info(_(f"Q Cluster {self.name} stopping."))
@@ -83,46 +81,46 @@ class Cluster:
self.stop() self.stop()
@property @property
def stat(self): def stat(self) -> Status:
if self.sentinel: if self.sentinel:
return Stat.get(pid=self.pid, cluster_id=self.cluster_id) return Stat.get(pid=self.pid, cluster_id=self.cluster_id)
return Status(pid=self.pid, cluster_id=self.cluster_id) return Status(pid=self.pid, cluster_id=self.cluster_id)
@property @property
def name(self): def name(self) -> str:
return humanize(self.cluster_id.hex) return humanize(self.cluster_id.hex)
@property @property
def is_starting(self): def is_starting(self) -> bool:
return self.stop_event and self.start_event and not self.start_event.is_set() return self.stop_event and self.start_event and not self.start_event.is_set()
@property @property
def is_running(self): def is_running(self) -> bool:
return self.stop_event and self.start_event and self.start_event.is_set() return self.stop_event and self.start_event and self.start_event.is_set()
@property @property
def is_stopping(self): def is_stopping(self) -> bool:
return ( return (
self.stop_event self.stop_event
and self.start_event and self.start_event
and self.start_event.is_set() and self.start_event.is_set()
and self.stop_event.is_set() and self.stop_event.is_set()
) )
@property @property
def has_stopped(self): def has_stopped(self) -> bool:
return self.start_event is None and self.stop_event is None and self.sentinel return self.start_event is None and self.stop_event is None and self.sentinel
class Sentinel: class Sentinel:
def __init__( def __init__(
self, self,
stop_event, stop_event,
start_event, start_event,
cluster_id, cluster_id,
broker=None, broker=None,
timeout=Conf.TIMEOUT, timeout=Conf.TIMEOUT,
start=True, start=True,
): ):
# Make sure we catch signals for the pool # Make sure we catch signals for the pool
signal.signal(signal.SIGINT, signal.SIG_IGN) signal.signal(signal.SIGINT, signal.SIG_IGN)
@@ -154,7 +152,7 @@ class Sentinel:
self.spawn_cluster() self.spawn_cluster()
self.guard() self.guard()
def status(self): def status(self) -> str:
if not self.start_event.is_set() and not self.stop_event.is_set(): if not self.start_event.is_set() and not self.stop_event.is_set():
return Conf.STARTING return Conf.STARTING
elif self.start_event.is_set() and not self.stop_event.is_set(): elif self.start_event.is_set() and not self.stop_event.is_set():
@@ -166,7 +164,7 @@ class Sentinel:
return Conf.STOPPING return Conf.STOPPING
return Conf.STOPPED return Conf.STOPPED
def spawn_process(self, target, *args): def spawn_process(self, target, *args) -> Process:
""" """
:type target: function or class :type target: function or class
""" """
@@ -179,7 +177,7 @@ class Sentinel:
p.start() p.start()
return p return p
def spawn_pusher(self): def spawn_pusher(self) -> Process:
return self.spawn_process(pusher, self.task_queue, self.event_out, self.broker) return self.spawn_process(pusher, self.task_queue, self.event_out, self.broker)
def spawn_worker(self): def spawn_worker(self):
@@ -187,7 +185,7 @@ class Sentinel:
worker, self.task_queue, self.result_queue, Value("f", -1), self.timeout worker, self.task_queue, self.result_queue, Value("f", -1), self.timeout
) )
def spawn_monitor(self): def spawn_monitor(self) -> Process:
return self.spawn_process(monitor, self.result_queue, self.broker) return self.spawn_process(monitor, self.result_queue, self.broker)
def reincarnate(self, process): def reincarnate(self, process):
@@ -310,9 +308,10 @@ class Sentinel:
Stat(self).save() Stat(self).save()
def pusher(task_queue, event, broker=None): def pusher(task_queue: Queue, event: Event, broker: Broker = None):
""" """
Pulls tasks of the broker and puts them in the task queue Pulls tasks of the broker and puts them in the task queue
:type broker:
:type task_queue: multiprocessing.Queue :type task_queue: multiprocessing.Queue
:type event: multiprocessing.Event :type event: multiprocessing.Event
""" """
@@ -545,9 +544,9 @@ def scheduler(broker=None):
try: try:
with db.transaction.atomic(using=Schedule.objects.db): with db.transaction.atomic(using=Schedule.objects.db):
for s in ( for s in (
Schedule.objects.select_for_update() Schedule.objects.select_for_update()
.exclude(repeats=0) .exclude(repeats=0)
.filter(next_run__lt=timezone.now()) .filter(next_run__lt=timezone.now())
): ):
args = () args = ()
kwargs = {} kwargs = {}
@@ -588,7 +587,11 @@ def scheduler(broker=None):
break break
# arrow always returns a tz aware datetime, and we don't want # arrow always returns a tz aware datetime, and we don't want
# this when we explicitly configured django with USE_TZ=False # this when we explicitly configured django with USE_TZ=False
s.next_run = next_run.datetime if settings.USE_TZ else next_run.datetime.replace(tzinfo=None) s.next_run = (
next_run.datetime
if settings.USE_TZ
else next_run.datetime.replace(tzinfo=None)
)
s.repeats += -1 s.repeats += -1
# send it to the cluster # send it to the cluster
q_options["broker"] = broker q_options["broker"] = broker
@@ -635,7 +638,7 @@ def close_old_django_connections():
db.close_old_connections() db.close_old_connections()
def set_cpu_affinity(n, process_ids, actual=not Conf.TESTING): def set_cpu_affinity(n: int, process_ids: list, actual: bool = not Conf.TESTING):
""" """
Sets the cpu affinity for the supplied processes. Sets the cpu affinity for the supplied processes.
Requires the optional psutil module. Requires the optional psutil module.
+18 -8
View File
@@ -2,8 +2,15 @@ import datetime
import time import time
import zlib import zlib
from django.core.signing import BadSignature, SignatureExpired, b64_decode, JSONSerializer, \ from django.core.signing import (
Signer as Sgnr, TimestampSigner as TsS, dumps BadSignature,
SignatureExpired,
b64_decode,
JSONSerializer,
Signer as Sgnr,
TimestampSigner as TsS,
dumps,
)
from django.utils import baseconv from django.utils import baseconv
from django.utils.crypto import constant_time_compare from django.utils.crypto import constant_time_compare
from django.utils.encoding import force_bytes, force_str, force_text from django.utils.encoding import force_bytes, force_str, force_text
@@ -16,7 +23,13 @@ The difference is that `this` loads function calls `TimestampSigner` and `Signer
""" """
def loads(s, key=None, salt='django.core.signing', serializer=JSONSerializer, max_age=None): def loads(
s,
key=None,
salt: str = "django.core.signing",
serializer=JSONSerializer,
max_age=None,
):
""" """
Reverse of dumps(), raise BadSignature if signature fails. Reverse of dumps(), raise BadSignature if signature fails.
@@ -26,7 +39,7 @@ def loads(s, key=None, salt='django.core.signing', serializer=JSONSerializer, ma
# operate on bytes. # operate on bytes.
base64d = force_bytes(TimestampSigner(key, salt=salt).unsign(s, max_age=max_age)) base64d = force_bytes(TimestampSigner(key, salt=salt).unsign(s, max_age=max_age))
decompress = False decompress = False
if base64d[:1] == b'.': if base64d[:1] == b".":
# It's compressed; uncompress it first # It's compressed; uncompress it first
base64d = base64d[1:] base64d = base64d[1:]
decompress = True decompress = True
@@ -37,7 +50,6 @@ def loads(s, key=None, salt='django.core.signing', serializer=JSONSerializer, ma
class Signer(Sgnr): class Signer(Sgnr):
def unsign(self, signed_value): def unsign(self, signed_value):
# force_str is removed in Django 2.0 # force_str is removed in Django 2.0
signed_value = force_str(signed_value) signed_value = force_str(signed_value)
@@ -57,7 +69,6 @@ calling `this` Signer.
class TimestampSigner(Signer, TsS): class TimestampSigner(Signer, TsS):
def unsign(self, value, max_age=None): def unsign(self, value, max_age=None):
""" """
Retrieve original value and check it wasn't signed more Retrieve original value and check it wasn't signed more
@@ -72,6 +83,5 @@ class TimestampSigner(Signer, TsS):
# Check timestamp is not older than max_age # Check timestamp is not older than max_age
age = time.time() - timestamp age = time.time() - timestamp
if age > max_age: if age > max_age:
raise SignatureExpired( raise SignatureExpired("Signature age %s > %s seconds" % (age, max_age))
'Signature age %s > %s seconds' % (age, max_age))
return value return value
+7 -6
View File
@@ -1,10 +1,9 @@
""" """
The code is derived from https://github.com/althonos/pronto/commit/3384010dfb4fc7c66a219f59276adef3288a886b The code is derived from https://github.com/althonos/pronto/commit/3384010dfb4fc7c66a219f59276adef3288a886b
""" """
import sys
import multiprocessing import multiprocessing
import multiprocessing.queues import multiprocessing.queues
import sys
class SharedCounter: class SharedCounter:
@@ -22,7 +21,7 @@ class SharedCounter:
""" """
def __init__(self, n=0): def __init__(self, n=0):
self.count = multiprocessing.Value('i', n) self.count = multiprocessing.Value("i", n)
def increment(self, n=1): def increment(self, n=1):
""" Increment the counter by n (default = 1) """ """ Increment the counter by n (default = 1) """
@@ -52,7 +51,9 @@ class Queue(multiprocessing.queues.Queue):
if sys.version_info < (3, 0): if sys.version_info < (3, 0):
super(Queue, self).__init__(*args, **kwargs) super(Queue, self).__init__(*args, **kwargs)
else: else:
super(Queue, self).__init__(*args, ctx=multiprocessing.get_context(), **kwargs) super(Queue, self).__init__(
*args, ctx=multiprocessing.get_context(), **kwargs
)
self.size = SharedCounter(0) self.size = SharedCounter(0)
@@ -65,10 +66,10 @@ class Queue(multiprocessing.queues.Queue):
self.size.increment(-1) self.size.increment(-1)
return x return x
def qsize(self): def qsize(self) -> int:
""" Reliable implementation of multiprocessing.Queue.qsize() """ """ Reliable implementation of multiprocessing.Queue.qsize() """
return self.size.value return self.size.value
def empty(self): def empty(self) -> bool:
""" Reliable implementation of multiprocessing.Queue.empty() """ """ Reliable implementation of multiprocessing.Queue.empty() """
return not self.qsize() > 0 return not self.qsize() > 0
+13 -18
View File
@@ -1,44 +1,39 @@
"""Package signing.""" """Package signing."""
try: import pickle
import cPickle as pickle
except ImportError:
import pickle
from django_q import core_signing as signing from django_q import core_signing as signing
from django_q.conf import Conf from django_q.conf import Conf
BadSignature = signing.BadSignature BadSignature = signing.BadSignature
class SignedPackage: class SignedPackage:
"""Wraps Django's signing module with custom Pickle serializer.""" """Wraps Django's signing module with custom Pickle serializer."""
@staticmethod @staticmethod
def dumps(obj, compressed=Conf.COMPRESSED): def dumps(obj, compressed=Conf.COMPRESSED):
return signing.dumps(obj, return signing.dumps(
key=Conf.SECRET_KEY, obj,
salt=Conf.PREFIX, key=Conf.SECRET_KEY,
compress=compressed, salt=Conf.PREFIX,
serializer=PickleSerializer) compress=compressed,
serializer=PickleSerializer,
)
@staticmethod @staticmethod
def loads(obj): def loads(obj):
return signing.loads(obj, return signing.loads(
key=Conf.SECRET_KEY, obj, key=Conf.SECRET_KEY, salt=Conf.PREFIX, serializer=PickleSerializer
salt=Conf.PREFIX, )
serializer=PickleSerializer)
class PickleSerializer: class PickleSerializer:
"""Simple wrapper around Pickle for signing.dumps and signing.loads.""" """Simple wrapper around Pickle for signing.dumps and signing.loads."""
@staticmethod @staticmethod
def dumps(obj): def dumps(obj) -> bytes:
return pickle.dumps(obj, protocol=pickle.HIGHEST_PROTOCOL) return pickle.dumps(obj, protocol=pickle.HIGHEST_PROTOCOL)
@staticmethod @staticmethod
def loads(data): def loads(data) -> any:
return pickle.loads(data) return pickle.loads(data)
+12 -7
View File
@@ -1,6 +1,9 @@
import socket import socket
from typing import Union
from django.utils import timezone from django.utils import timezone
from django_q.brokers import get_broker
from django_q.brokers import get_broker, Broker
from django_q.conf import Conf, logger from django_q.conf import Conf, logger
from django_q.signing import SignedPackage, BadSignature from django_q.signing import SignedPackage, BadSignature
@@ -47,18 +50,18 @@ class Stat(Status):
self.pusher = sentinel.pusher.pid self.pusher = sentinel.pusher.pid
self.workers = [w.pid for w in sentinel.pool] self.workers = [w.pid for w in sentinel.pool]
def uptime(self): def uptime(self) -> float:
return (timezone.now() - self.tob).total_seconds() return (timezone.now() - self.tob).total_seconds()
@property @property
def key(self): def key(self) -> str:
""" """
:return: redis key for this cluster statistic :return: redis key for this cluster statistic
""" """
return self.get_key(self.cluster_id) return self.get_key(self.cluster_id)
@staticmethod @staticmethod
def get_key(cluster_id): def get_key(cluster_id) -> str:
""" """
:param cluster_id: cluster ID :param cluster_id: cluster ID
:return: redis key for the cluster statistic :return: redis key for the cluster statistic
@@ -71,13 +74,15 @@ class Stat(Status):
except Exception as e: except Exception as e:
logger.error(e) logger.error(e)
def empty_queues(self): def empty_queues(self) -> bool:
return self.done_q_size + self.task_q_size == 0 return self.done_q_size + self.task_q_size == 0
@staticmethod @staticmethod
def get(pid, cluster_id, broker=None): def get(pid: int, cluster_id: str, broker: Broker = None) -> Union[Status, None]:
""" """
gets the current status for the cluster gets the current status for the cluster
:param pid:
:param broker: an optional broker instance
:param cluster_id: id of the cluster :param cluster_id: id of the cluster
:return: Stat or Status :return: Stat or Status
""" """
@@ -92,7 +97,7 @@ class Stat(Status):
return Status(pid=pid, cluster_id=cluster_id) return Status(pid=pid, cluster_id=cluster_id)
@staticmethod @staticmethod
def get_all(broker=None): def get_all(broker: Broker = None)->list:
""" """
Get the status for all currently running clusters with the same prefix Get the status for all currently running clusters with the same prefix
and secret key. and secret key.