import time,uuid,threading,logging
import zmq
from SimpleWebSocketServer import SimpleWebSocketServer, WebSocket

'''
pubsubws bridge

all ws connection cant share same socket, but can share context

'''
#yes global threadsafe variable
#but they are used readonly 
#like constant
logger = logging.getLogger(__name__)
Done = False    #lazy exit signal for thread
ctx = zmq.Context()
################################################websocket
class wshandler(WebSocket):
    def run(self):
        '''
        thread that receive from SUB socket and forward to respective websocket or all if uuid='*'
        sub inproc://subtows
        '''
        logger.info('subscriber to websocket bridge online')
        zsub = ctx.socket(zmq.SUB)
        zsub.connect('inproc://subtows')
        zsub.setsockopt(zmq.SUBSCRIBE,self.uuid)
        zsub.setsockopt(zmq.SUBSCRIBE,'*')
        while self.connected:
            uuid,msg = zsub.recv_multipart()
            self.sendMessage(unicode(msg,'utf-8'))
            logger.debug('subscriber to websocket bridge forwared %s:%s...'%(uuid,msg[197:]))
        zsub.close()
        logger.info('subscriber to websocket bridge offline')

    def handleMessage(self):
        '''
        will send data to pub with uuid
        '''
        logger.debug('websocket got %s:%s, redirect to pub'%(self.uuid,self.data[197:]+'...'))
        self.zpub.send_multipart((self.uuid,str(self.data)))

    def handleConnected(self):
        '''
        generate uuid
        pub to inproc://wstopub
        create pub socket
        '''
        self.uuid = str(uuid.uuid4())
        logger.debug('websocket %s connect'%self.uuid)
        self.connected = True
        self.t = threading.Thread(target=self.run,args=())
        self.t.daemon = True
        self.t.start()
        self.zpub = ctx.socket(zmq.PUB)
        self.zpub.connect('inproc://wstopub')

    def handleClose(self):
        '''
        stop thread and join
        close pub connection
        '''
        logger.debug('websocket %s disconnect'%self.uuid)
        self.connected=False
        self.zpub.close()

###########################################################################
#TODO turn these thread into a zmq specialized solution
#http://api.zeromq.org/3-2:zmq-proxy
#http://pyzmq.readthedocs.io/en/latest/api/zmq.devices.html
#http://learning-0mq-with-pyzmq.readthedocs.io/en/latest/pyzmq/devices/forwarder.html
#this server will be easyer to port to C
def statz(stats=None,payload=''):
    if not stats:
        return (0,0,0)
    sum, cnt, avg = stats
    sum = sum + len(payload)
    cnt = cnt + 1
    avg = 1.0*sum/cnt
    return (sum,cnt,avg)

def wstopub(port=7777):
    '''
    outgoing ws to pub socket proxy
    '''
    pub = ctx.socket(zmq.PUB)
    pub.bind('tcp://*:%s'%port)
    sub = ctx.socket(zmq.SUB)
    sub.bind('inproc://wstopub')
    sub.setsockopt(zmq.SUBSCRIBE,'')
    stats = statz()
    logger.info('publisher bind %s'%port)
    while not Done:
        uuid,msg = sub.recv_multipart()
        if msg:
            pub.send_multipart((uuid,msg))
            stats = statz(stats,msg)
            #logger.debug('publisher to ws %s:%s... avg/msg:%s'%(uuid,msg[197:],stats[2]))
    logger.info('publisher unbind')
    pub.close()
    sub.close()

def subtows(port=6666):
    '''
    incoming sub socket to ws proxy
    '''
    sub = ctx.socket(zmq.SUB)
    sub.bind('tcp://*:%s'%port)
    sub.setsockopt(zmq.SUBSCRIBE,'')
    pub = ctx.socket(zmq.PUB)
    pub.bind('inproc://subtows')
    stats = statz()
    logger.info('subscriber bind %s'%port)
    while not Done:
        uuid,msg = sub.recv_multipart()
        if msg:
            pub.send_multipart((uuid,msg))
            stats = statz(stats,msg)
            #logger.debug('subscriber from ws %s:%s... avg/msg:%s'%(uuid,msg[197:],stats[2]))
    logger.info('subscriber unbind')
    pub.close()
    sub.close()
###########################################################################################
def testThread(inport=7777,outport=6666):
    time.sleep(2)
    print('begin test')
    _ctx = zmq.Context()
    sub = _ctx.socket(zmq.SUB)
    sub.connect('tcp://127.0.0.1:%s'%inport)
    sub.setsockopt(zmq.SUBSCRIBE,'')
    pub = _ctx.socket(zmq.PUB)
    pub.connect('tcp://127.0.0.1:%s'%outport)
    while not Done:
        uuid,msg = sub.recv_multipart()
        print('test recv',msg,'from',uuid)
        print('resend to all')
        pub.send_multipart(('*',msg))
    sub.close()
    pub.close()
    _ctx.term()
################################################################################

if __name__ == '__main__':
    allLogLvl = {
        'debug':logging.DEBUG,
        'info':logging.INFO,
        'warning':logging.WARNING,
        'error':logging.ERROR,
        'critical':logging.CRITICAL
    }
    import argparse
    parser = argparse.ArgumentParser(description='websocket <=> zeromq PubSub bridge')
    parser.add_argument('--pub',type=int,default=7777,help="tcp port to bind for the publisher")
    parser.add_argument('--sub',type=int,default=6666,help="tcp port to bind for the reverse subscriber")
    parser.add_argument('--ws',type=int,default=8888,help="http port to bind for the websocket server")
    parser.add_argument('--log',default='info',choices=allLogLvl,help="logging level")
    args = parser.parse_args()
    #logger.setLevel(allLogLvl[args.log])
    #logger.addHandler(logging.StreamHandler())
    logging.basicConfig(level=allLogLvl[args.log])


    thlst = [
        #threading.Thread(target=testThread,args=(7777,6666)),
        threading.Thread(target=subtows, args=(args.sub,)),
        threading.Thread(target=wstopub, args=(args.pub,))
    ]
    for th in thlst:
        th.daemon = True
        th.start()
    server = SimpleWebSocketServer('', args.ws, wshandler)
    server.serveforever()

    #ok there are mutation in so called constant
    #but its in the same scope/thread
    Done = True
    ctx.term()
    print('Done')
