diff options
Diffstat (limited to 'joblib.py')
| -rw-r--r-- | joblib.py | 246 |
1 files changed, 62 insertions, 184 deletions
@@ -1,206 +1,84 @@ #!/usr/bin/python3 -import asyncio -import asyncpg import json import logging +import pika class JobQueue(object): def __init__(self): - self._con = None - self._listens = {} + self._qcon = None + self._qchan = None - def _listener(self, con, pid, - channel, payload): - """Fire the registered callback. Must not be async.""" - self._listens[channel](payload) + def connect(self, **args): + """Connect to RabbitMQ.""" + logging.info("connecting to broker") + self._qcon = pika.BlockingConnection(pika.ConnectionParameters( + **args)) + self._qchan = self._qcon.channel() - async def _hook_listener(self): - """Register callbacks.""" - if len(self._listens.keys()) > 0: - for channel in self._listens.keys(): - await self._con.add_listener(channel, - self._listener) - - async def connect(self, **args): - """Connect the database.""" - logging.info("connecting to queue database") - while True: - try: - self._con = await asyncpg.connect(**args) - await self._hook_listener() - break - except asyncpg.exceptions.CannotConnectNowError: - logging.warning( - "queue database is not ready yet, retrying") - await asyncio.sleep(1) - except asyncpg.exceptions.TooManyConnectionsError: - logging.warning( - "queue database has too many clients, retrying") - await asyncio.sleep(1) - - async def setup(self): - """Initialize the database.""" - await self._con.execute(""" - CREATE TABLE IF NOT EXISTS - jobs ( - id SERIAL PRIMARY KEY, - priority INT DEFAULT 0, - status TEXT DEFAULT 'open', - type TEXT, - payload JSONB - ) - """) - await self._con.execute(""" - CREATE UNIQUE INDEX - IF NOT EXISTS - jobs_priority_id_idx - ON jobs(priority, id) - """) - - async def listen(self, channel, callback): - """Hook a notification listener.""" - self._listens[channel] = callback - - async def request(self, jobtype, payload, priority=0): + def request(self, queue, payload): """Create a new job and return its id.""" - if isinstance(payload, list): - payloads = [json.dumps(p) - for p in payload] - else: - payloads = [json.dumps(payload)] - if isinstance(priority, list): - priorities = priority - else: - priorities = [priority] * len(payloads) - insert = await self._con.prepare(""" - INSERT INTO jobs(type, payload, priority) - VALUES($1, $2, $3) - RETURNING id - """) - ids = [] - async with self._con.transaction(): - for pl, pr in zip(payloads, priorities): - ids.append(await insert.fetchval(jobtype, pl, pr)) - await self._con.execute("SELECT pg_notify($1 || '_open', '')", - jobtype) + body = json.dumps(payload) + self._qchan.queue_declare(queue=queue, durable=True) + self._qchan.basic_publish(exchange="", + routing_key=queue, + body=body, + properties=pika.BasicProperties( + delivery_mode = 2 + )) - if isinstance(payload, list): - return ids - else: - return ids[0] +class JobFailed(Exception): + pass - async def acquire(self, jobtype, length=None): - """Mark a job as running, return id, payload and priority. - Return (None, None, None) if no job is available.""" - if length is None: - limit = 1 - else: - limit = length - while True: - try: - # do not allow async access - async with self._con.transaction(isolation="serializable"): - result = await self._con.fetch(""" - UPDATE jobs SET STATUS='running' - FROM ( - SELECT id FROM jobs - WHERE status='open' AND type=$1 - ORDER BY priority, id - LIMIT $2 - ) AS open_jobs - WHERE jobs.id=open_jobs.id - RETURNING jobs.id, jobs.payload, jobs.priority - """, jobtype, limit) - if len(result) == 0 and length is None: - # no jobs available - # backwards compatibility - return None, None, None - jobs = [(r[0], json.loads(r[1]), r[2]) for r in result] +class Worker(JobQueue): + """Abstract service worker class.""" + def __init__(self, jobtype): + super().__init__() + self._queue = jobtype + self._up = False + self._notifq = [] - await self._con.execute("SELECT pg_notify($1 || '_running', '')", - jobtype) + def setup(self): + # override + pass - if length is None: - return jobs[0] - else: - return jobs - except asyncpg.exceptions.SerializationError: - # job is being picked up by another worker, try again - pass + def work(self, payload): + # override + pass - async def status(self, jobid): - """Return the status of a job.""" - return await self._con.fetchval(""" - SELECT status - FROM jobs WHERE - id=$1 - """, jobid) + def commit(self, failed): + # override + pass - async def finish(self, jobid, jobtype): - """Mark jobs as completed.""" - if not isinstance(jobid, list): - jobids = [(jobid,)] + def request(self, queue, payload, now=True): + if now: + super().request(queue, payload) else: - jobids = [(jid,) for jid in jobid] - async with self._con.transaction(): - await self._con.executemany(""" - UPDATE jobs SET status='finished' - WHERE id=$1 - """, jobids) - await self._con.execute(""" - SELECT pg_notify($1 || '_finished', '') - """, jobtype) + # send after COMMIT + self._notifq.append((queue, payload)) - async def fail(self, jobid, jobtype, reason): - """Mark a job as failed.""" - if not isinstance(jobid, list): - jobids = [jobid] - else: - jobids = jobid - if not isinstance(reason, list): - reasons = [json.dumps({"error": reason})] - else: - reasons = [json.dumps({"error": r}) - for r in reason] - assert len(jobids) == len(reasons) - async with self._con.transaction(): - await self._con.executemany(""" - UPDATE jobs SET status='failed', - payload=payload||$2::jsonb - WHERE id=$1 - """, zip(jobids, reasons)) - await self._con.execute(""" - SELECT pg_notify($1 || '_failed', '') - """, jobtype) - - async def reset(self, jobid, jobtype): - """Mark a job as open.""" - if not isinstance(jobid, list): - jobids = [(jobid,)] - else: - jobids = [(jid,) for jid in jobid] - async with self._con.transaction(): - await self._con.executemany(""" - UPDATE jobs SET status='open' - WHERE id=$1 - """, jobids) - await self._con.execute(""" - SELECT pg_notify($1 || '_open', '') - """, jobtype) + def run(self): + self._qchan.queue_declare(queue=self._queue, durable=True) + self._qchan.basic_qos(prefetch_count=1) + for msg in self._qchan.consume(queue=self._queue, + inactivity_timeout=1): + # TODO force commit after n, c+=1 + if msg is None: + # idling + self.commit(False) + # send notifs dependant on COMMIT + for notif in self._notifq: + self.request(*notif) + self._notifq = [] + continue - async def cleanup(self): - """Reopen all unfinished jobs.""" - while True: + method, properties, body = msg + # run a job + payload = json.loads(body) try: - async with self._con.transaction(isolation="serializable"): - await self._con.execute(""" - UPDATE jobs - SET status='open' - WHERE status='running' - """) - return - except asyncpg.exceptions.SerializationError: - pass + self.work(payload) + except JobFailed as err: + pass # TODO + self._qchan.basic_ack(delivery_tag=method.delivery_tag) |
