summaryrefslogtreecommitdiff
path: root/api.py
blob: 49a2b0d63777ebb609f8d21adf8ff13e94a8a32f (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
#!/usr/bin/python

import os
import asyncio
import logging

import joblib.worker
import classifier


#tf.logging.set_verbosity(tf.logging.WARNING)

queue_db = {
    "host": os.environ.get("POSTGRESQL_SOURCE_HOST") or "localhost",
    "port": os.environ.get("POSTGRESQL_SOURCE_PORT") or 5433,
    "user": os.environ.get("POSTGRESQL_SOURCE_USER") or "vainraw",
    "password": os.environ.get("POSTGRESQL_SOURCE_PASSWORD") or "vainraw",
    "database": os.environ.get("POSTGRESQL_SOURCE_DB") or "vainsocial-raw"
}

db_config = {
    "host": os.environ.get("POSTGRESQL_DEST_HOST") or "localhost",
    "port": os.environ.get("POSTGRESQL_DEST_PORT") or 5432,
    "user": os.environ.get("POSTGRESQL_DEST_USER") or "vainweb",
    "password": os.environ.get("POSTGRESQL_DEST_PASSWORD") or "vainweb",
    "database": os.environ.get("POSTGRESQL_DEST_DB") or "vainsocial-web"
}


class Analyzer(joblib.worker.Worker):
    def __init__(self):
        self._queries = {}
        super().__init__(jobtype="analyze")
        self.classifiers = []

    async def connect(self, dbconf, queuedb):
        """Connect to database."""
        logging.warning("connecting to database")
        await super().connect(**queuedb)
        cl = classifier.KDAClassifier()
        cl.connect(**dbconf)
        self.classifiers.append(cl)
#        cl = classifier.RoleClassifier()
#        cl.connect(**dbconf)
#        self.classifiers.append(cl)

    async def setup(self):
        """Setup the model."""
        for cl in self.classifiers:
            cl.train()

    async def _windup(self):
        for cl in self.classifiers:
            cl.windup()
        self._participants = []

    async def _teardown(self, failed):
        if len(self._participants) > 0:
            # TODO if this fails, job is still marked as finished
            for cl in self.classifiers:
                cl.classify_db(self._participants)
        for cl in self.classifiers:
            cl.teardown()
    
    async def _execute_job(self, jobid, payload, priority):
        object_id = payload["id"]
        object_type = payload["type"]
        if object_type != "participant":
            return
        self._participants.append(object_id)
        logging.info("%s: classifying '%s', %s", jobid,
                     object_type, object_id)

async def startup():
    for _ in range(1):
        worker = Analyzer()
        await worker.connect(db_config, queue_db)
        await worker.setup()
        await worker.start(batchlimit=10000)


logging.basicConfig(level=logging.DEBUG)

loop = asyncio.get_event_loop()
loop.run_until_complete(startup())
loop.run_forever()