summaryrefslogtreecommitdiff
path: root/joblib.py
diff options
context:
space:
mode:
Diffstat (limited to 'joblib.py')
-rw-r--r--joblib.py246
1 files changed, 62 insertions, 184 deletions
diff --git a/joblib.py b/joblib.py
index 04a37d0..4cc3b94 100644
--- a/joblib.py
+++ b/joblib.py
@@ -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)