mirror of
https://github.com/wassname/catalyst.git
synced 2026-08-12 11:50:11 +08:00
Merge branch 'master' of github.com:quantopian/zipline
This commit is contained in:
@@ -1,3 +1,5 @@
|
||||
#zeromq related
|
||||
pyzmq==2.1.11
|
||||
gevent-zeromq==0.2.2
|
||||
msgpack-python==0.1.12
|
||||
humanhash==0.0.1
|
||||
|
||||
+167
-56
@@ -1,10 +1,13 @@
|
||||
"""
|
||||
Commonly used messaging components.
|
||||
"""
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
import datetime
|
||||
import socket
|
||||
import humanhash
|
||||
|
||||
import zipline.util as qutil
|
||||
from zipline.protocol import CONTROL_PROTOCOL, COMPONENT_STATE
|
||||
|
||||
class Component(object):
|
||||
|
||||
@@ -14,8 +17,6 @@ class Component(object):
|
||||
|
||||
- sync_address: socket address used for synchronizing the start of all workers, heartbeating, and exit notification
|
||||
will be used in REP/REQ sockets. Bind is always on the REP side.
|
||||
- control_address: socket address used for controlling and
|
||||
monitoring the status of the simulation
|
||||
- data_address: socket address used for data sources to stream their records.
|
||||
will be used in PUSH/PULL sockets between data sources and a ParallelBuffer (aka the Feed). Bind
|
||||
will always be on the PULL side (we always have N producers and 1 consumer)
|
||||
@@ -31,17 +32,23 @@ class Component(object):
|
||||
will also return a Poller.
|
||||
|
||||
"""
|
||||
self.zmq = None
|
||||
self.context = None
|
||||
self.addresses = None
|
||||
self.out_socket = None
|
||||
self.gevent_needed = False
|
||||
self.killed = False
|
||||
self.zmq = None
|
||||
self.context = None
|
||||
self.addresses = None
|
||||
self.out_socket = None
|
||||
self.gevent_needed = False
|
||||
self.killed = False
|
||||
self.heartbeat_timeout = 2000
|
||||
self.state_flag = COMPONENT_STATE.OK # OK | DONE | EXCEPTION
|
||||
|
||||
# TODO: could probably mkae this into a property instead of a
|
||||
# method
|
||||
def get_id(self):
|
||||
raise NotImplementedError
|
||||
# Humanhashes make this way easier to debug because they
|
||||
# stick in your mind unlike a 32 byte string of random hex.
|
||||
self.guid = uuid.uuid4()
|
||||
self.huid = humanhash.humanize(self.guid.hex)
|
||||
|
||||
# ------------
|
||||
# Core Methods
|
||||
# ------------
|
||||
|
||||
def open(self):
|
||||
raise NotImplementedError
|
||||
@@ -62,17 +69,12 @@ class Component(object):
|
||||
def do_work(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def run(self):
|
||||
|
||||
fail = None
|
||||
|
||||
#try:
|
||||
#TODO: can't initialize these values in the __init__?
|
||||
self.done = False
|
||||
def _run(self):
|
||||
self.done = False # TODO: use state flag
|
||||
self.sockets = []
|
||||
|
||||
if self.gevent_needed:
|
||||
qutil.LOGGER.info("Loading gevent specific zmq for {id}".format(id=self.get_id()))
|
||||
qutil.LOGGER.info("Loading gevent specific zmq for {id}".format(id=self.get_id))
|
||||
import gevent_zeromq
|
||||
self.zmq = gevent_zeromq.zmq
|
||||
else:
|
||||
@@ -80,6 +82,8 @@ class Component(object):
|
||||
self.zmq = zmq
|
||||
|
||||
self.context = self.zmq.Context()
|
||||
self.setup_poller()
|
||||
|
||||
self.open()
|
||||
self.setup_sync()
|
||||
self.setup_control()
|
||||
@@ -89,50 +93,102 @@ class Component(object):
|
||||
for sock in self.sockets:
|
||||
sock.close()
|
||||
|
||||
#except Exception as e:
|
||||
#qutil.LOGGER.exception("Unexpected error in run for {id}.".format(id=self.get_id()))
|
||||
#fail = e
|
||||
def run(self, catch_exceptions=False):
|
||||
"""
|
||||
Run the component.
|
||||
|
||||
#finally:
|
||||
Optionally takes an argument to catch and log all exceptions raised
|
||||
during execution ues this with care since it makes it very hard to
|
||||
debug since it mucks up your stacktraces.
|
||||
"""
|
||||
|
||||
#if(self.context != None):
|
||||
#self.context.destroy()
|
||||
fail = None
|
||||
|
||||
#if fail:
|
||||
#raise fail
|
||||
if catch_exceptions:
|
||||
try:
|
||||
self._run()
|
||||
except Exception as exc:
|
||||
# TODO, we want to do this grab the stack
|
||||
# frame so we can preserve stacktraces when we
|
||||
# reraise the exception.
|
||||
self.signal_exception(exc)
|
||||
fail = exc
|
||||
finally:
|
||||
if self.context:
|
||||
self.context.destroy()
|
||||
if fail:
|
||||
raise fail
|
||||
else:
|
||||
self._run()
|
||||
if(self.context != None):
|
||||
self.context.destroy()
|
||||
|
||||
def loop(self):
|
||||
while not self.done:
|
||||
"""
|
||||
Loop to do work while we still have work to do.
|
||||
"""
|
||||
while not self.done: # TODO: use state flag
|
||||
self.confirm()
|
||||
self.do_work()
|
||||
|
||||
def confirm(self):
|
||||
"""
|
||||
Send a synchronization request to the host.
|
||||
"""
|
||||
|
||||
# TODO: proper framing
|
||||
self.sync_socket.send(self.get_id + ":RUN")
|
||||
|
||||
self.receive_sync_ack() # blocking
|
||||
|
||||
# ----------------------
|
||||
# Internal Maintenance
|
||||
# ----------------------
|
||||
|
||||
def signal_exception(self, exc=None):
|
||||
self.state_flag = COMPONENT_STATE.EXCEPTION
|
||||
qutil.LOGGER.exception("Unexpected error in run for {id}.".format(id=self.get_id))
|
||||
|
||||
def signal_done(self):
|
||||
#notify down stream components that we're done
|
||||
if(self.out_socket != None):
|
||||
self.out_socket.send("DONE")
|
||||
"""
|
||||
Notify down stream components that we're done.
|
||||
"""
|
||||
|
||||
self.state_flag = COMPONENT_STATE.DONE
|
||||
|
||||
if self.out_socket:
|
||||
self.out_socket.send(str(CONTROL_PROTOCOL.DONE))
|
||||
|
||||
#notify host we're done
|
||||
self.sync_socket.send(self.get_id() + ":DONE")
|
||||
# TODO: proper framing
|
||||
self.sync_socket.send(self.get_id + ":" + str(CONTROL_PROTOCOL.DONE))
|
||||
|
||||
self.receive_sync_ack()
|
||||
#notify internal work look that we're done
|
||||
self.done = True
|
||||
self.done = True # TODO: use state flag
|
||||
|
||||
# TODO: probably don't need a method here ... or move into
|
||||
# higher level framing protocol
|
||||
def is_done_message(self, message):
|
||||
return message == "DONE"
|
||||
# -----------
|
||||
# Messaging
|
||||
# -----------
|
||||
|
||||
def confirm(self):
|
||||
# send a synchronization request to the host
|
||||
self.sync_socket.send(self.get_id() + ":RUN")
|
||||
self.receive_sync_ack()
|
||||
def setup_poller(self):
|
||||
"""
|
||||
Setup the poller used for multiplexing the incoming data
|
||||
handling sockets.
|
||||
"""
|
||||
|
||||
self.poll = self.zmq.Poller()
|
||||
|
||||
def receive_sync_ack(self):
|
||||
# wait for synchronization reply from the host
|
||||
socks = dict(self.sync_poller.poll(2000)) #timeout after 2 seconds.
|
||||
"""
|
||||
Wait for synchronization reply from the host.
|
||||
"""
|
||||
|
||||
socks = dict(self.sync_poller.poll(self.heartbeat_timeout))
|
||||
if self.sync_socket in socks and socks[self.sync_socket] == self.zmq.POLLIN:
|
||||
message = self.sync_socket.recv()
|
||||
else:
|
||||
raise Exception("Sync ack timed out on response for {id}".format(id=self.get_id()))
|
||||
raise Exception("Sync ack timed out on response for {id}".format(id=self.get_id))
|
||||
|
||||
def bind_data(self):
|
||||
return self.bind_pull_socket(self.addresses['data_address'])
|
||||
@@ -161,10 +217,11 @@ class Component(object):
|
||||
def bind_pull_socket(self, addr):
|
||||
pull_socket = self.context.socket(self.zmq.PULL)
|
||||
pull_socket.bind(addr)
|
||||
poller = self.zmq.Poller()
|
||||
poller.register(pull_socket, self.zmq.POLLIN)
|
||||
self.poll.register(pull_socket, self.zmq.POLLIN)
|
||||
|
||||
self.sockets.append(pull_socket)
|
||||
return pull_socket, poller
|
||||
|
||||
return pull_socket
|
||||
|
||||
def connect_push_socket(self, addr):
|
||||
push_socket = self.context.socket(self.zmq.PUSH)
|
||||
@@ -172,6 +229,7 @@ class Component(object):
|
||||
#push_socket.setsockopt(self.zmq.LINGER,0)
|
||||
self.sockets.append(push_socket)
|
||||
self.out_socket = push_socket
|
||||
|
||||
return push_socket
|
||||
|
||||
def bind_pub_socket(self, addr):
|
||||
@@ -179,32 +237,85 @@ class Component(object):
|
||||
pub_socket.bind(addr)
|
||||
#pub_socket.setsockopt(self.zmq.LINGER,0)
|
||||
self.out_socket = pub_socket
|
||||
|
||||
return pub_socket
|
||||
|
||||
def connect_sub_socket(self, addr):
|
||||
sub_socket = self.context.socket(self.zmq.SUB)
|
||||
sub_socket.connect(addr)
|
||||
sub_socket.setsockopt(self.zmq.SUBSCRIBE,'')
|
||||
poller = self.zmq.Poller()
|
||||
poller.register(sub_socket, self.zmq.POLLIN)
|
||||
self.sockets.append(sub_socket)
|
||||
return sub_socket, poller
|
||||
|
||||
self.poll.register(sub_socket, self.zmq.POLLIN)
|
||||
|
||||
return sub_socket
|
||||
|
||||
def setup_control(self):
|
||||
"""
|
||||
Set up the control socket. Used to monitor the the
|
||||
Set up the control socket. Used to monitor the
|
||||
overall status of the simulation and to forcefully tear
|
||||
down the simulation in case of a failure.
|
||||
"""
|
||||
pass
|
||||
assert self.controller
|
||||
|
||||
self.control_out = self.controller.message_sender()
|
||||
self.control_in = self.controller.message_listener()
|
||||
|
||||
self.poll.register(self.control_in, self.zmq.POLLIN)
|
||||
self.sockets.extend([self.control_in, self.control_out])
|
||||
|
||||
def setup_sync(self):
|
||||
qutil.LOGGER.debug("Connecting sync client for {id}".format(id=self.get_id()))
|
||||
"""
|
||||
Setup the sync socket and poller.
|
||||
"""
|
||||
|
||||
qutil.LOGGER.debug("Connecting sync client for {id}".format(id=self.get_id))
|
||||
|
||||
self.sync_socket = self.context.socket(self.zmq.REQ)
|
||||
self.sync_socket.connect(self.addresses['sync_address'])
|
||||
#self.sync_socket.setsockopt(self.zmq.LINGER,0)
|
||||
|
||||
# Explictly, a different poller for obvious reasons.
|
||||
# I'm not fond of having this poller init'd as a side
|
||||
# effect of a method call. Still thinking about where to
|
||||
# put it at the moment though...
|
||||
self.sync_poller = self.zmq.Poller()
|
||||
self.sync_poller.register(self.sync_socket, self.zmq.POLLIN)
|
||||
|
||||
self.sockets.append(self.sync_socket)
|
||||
|
||||
# ---------------------
|
||||
# Description and Debug
|
||||
# ---------------------
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return 'UNKNOWN COMPONENT'
|
||||
|
||||
def debug(self):
|
||||
"""
|
||||
Debug information about the component.
|
||||
"""
|
||||
return (
|
||||
self.get_id ,
|
||||
self.huid ,
|
||||
socket.gethostname() ,
|
||||
os.getpid() ,
|
||||
hex(id(self)) ,
|
||||
self.sockets ,
|
||||
)
|
||||
|
||||
def __repr__(self):
|
||||
"""
|
||||
Return a usefull string representation of the component
|
||||
to indicate its type, unique identifier, and computational
|
||||
context identifier name.
|
||||
"""
|
||||
|
||||
return "<{name} {uuid} at {host} {pid} {pointer}>".format(
|
||||
name = self.get_id ,
|
||||
uuid = self.huid ,
|
||||
host = socket.gethostname() ,
|
||||
pid = os.getpid() ,
|
||||
pointer = hex(id(self)) ,
|
||||
)
|
||||
|
||||
+57
-31
@@ -7,23 +7,27 @@ import datetime
|
||||
import zipline.util as qutil
|
||||
from zipline.component import Component
|
||||
|
||||
from zipline.protocol import CONTROL_PROTOCOL
|
||||
|
||||
class ComponentHost(Component):
|
||||
"""
|
||||
Components that can launch multiple sub-components, synchronize their start, and then wait for all
|
||||
components to be finished.
|
||||
Components that can launch multiple sub-components, synchronize their
|
||||
start, and then wait for all components to be finished.
|
||||
"""
|
||||
|
||||
def __init__(self, addresses, gevent_needed=False):
|
||||
Component.__init__(self)
|
||||
self.addresses = addresses
|
||||
|
||||
#workaround for defect in threaded use of strptime: http://bugs.python.org/issue11108
|
||||
# workaround for defect in threaded use of strptime:
|
||||
# http://bugs.python.org/issue11108
|
||||
qutil.parse_date("2012/02/13-10:04:28.114")
|
||||
|
||||
self.components = {}
|
||||
self.sync_register = {}
|
||||
self.timeout = datetime.timedelta(seconds=5)
|
||||
self.gevent_needed = gevent_needed
|
||||
self.heartbeat_timeout = 2000
|
||||
|
||||
self.feed = ParallelBuffer()
|
||||
self.merge = MergedParallelBuffer()
|
||||
@@ -47,13 +51,13 @@ class ComponentHost(Component):
|
||||
if self.controller:
|
||||
component.controller = self.controller
|
||||
|
||||
self.components[component.get_id()] = component
|
||||
self.sync_register[component.get_id()] = datetime.datetime.utcnow()
|
||||
self.components[component.get_id] = component
|
||||
self.sync_register[component.get_id] = datetime.datetime.utcnow()
|
||||
|
||||
if(isinstance(component, DataSource)):
|
||||
self.feed.add_source(component.get_id())
|
||||
self.feed.add_source(component.get_id)
|
||||
if(isinstance(component, BaseTransform)):
|
||||
self.merge.add_source(component.get_id())
|
||||
self.merge.add_source(component.get_id)
|
||||
|
||||
def unregister_component(self, component_id):
|
||||
del self.components[component_id]
|
||||
@@ -61,15 +65,20 @@ class ComponentHost(Component):
|
||||
|
||||
def setup_sync(self):
|
||||
"""
|
||||
Start the sync server.
|
||||
"""
|
||||
qutil.LOGGER.debug("Connecting sync server.")
|
||||
|
||||
self.sync_socket = self.context.socket(self.zmq.REP)
|
||||
self.sync_socket.bind(self.addresses['sync_address'])
|
||||
|
||||
self.poller = self.zmq.Poller()
|
||||
self.poller.register(self.sync_socket, self.zmq.POLLIN)
|
||||
# There is a namespace collision between three classes
|
||||
# which use the self.poller property to mean different
|
||||
# things.
|
||||
# =====================================================
|
||||
self.sync_poller = self.zmq.Poller()
|
||||
self.sync_poller.register(self.sync_socket, self.zmq.POLLIN)
|
||||
# =====================================================
|
||||
|
||||
self.sockets.append(self.sync_socket)
|
||||
|
||||
def open(self):
|
||||
@@ -83,29 +92,39 @@ class ComponentHost(Component):
|
||||
if len(self.components) == 0:
|
||||
qutil.LOGGER.info("Component register is empty.")
|
||||
return True
|
||||
|
||||
for source, last_dt in self.sync_register.iteritems():
|
||||
if((cur_time - last_dt) > self.timeout):
|
||||
qutil.LOGGER.info("Time out for {source}. Current component registery: {reg}".format(source=source, reg=self.components))
|
||||
if (cur_time - last_dt) > self.timeout:
|
||||
qutil.LOGGER.info(
|
||||
"Time out for {source}. Current component registery: {reg}".
|
||||
format(source=source, reg=self.components)
|
||||
)
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
def loop(self):
|
||||
|
||||
while not self.is_timed_out():
|
||||
# wait for synchronization request
|
||||
socks = dict(self.poller.poll(2000)) #timeout after 2 seconds.
|
||||
socks = dict(self.sync_poller.poll(self.heartbeat_timeout)) #timeout after 2 seconds.
|
||||
|
||||
if self.sync_socket in socks and socks[self.sync_socket] == self.zmq.POLLIN:
|
||||
msg = self.sync_socket.recv()
|
||||
parts = msg.split(':')
|
||||
if(len(parts) < 2):
|
||||
|
||||
if len(parts) != 2:
|
||||
qutil.LOGGER.info("got bad confirm: {msg}".format(msg=msg))
|
||||
sync_id = parts[0]
|
||||
status = parts[1]
|
||||
if(self.is_done_message(status)):
|
||||
continue
|
||||
|
||||
sync_id, status = parts
|
||||
|
||||
if status == str(CONTROL_PROTOCOL.DONE): # TODO: other way around
|
||||
qutil.LOGGER.info("{id} is DONE".format(id=sync_id))
|
||||
self.unregister_component(sync_id)
|
||||
else:
|
||||
self.sync_register[sync_id] = datetime.datetime.utcnow()
|
||||
|
||||
#qutil.LOGGER.info("confirmed {id}".format(id=msg))
|
||||
# send synchronization reply
|
||||
self.sync_socket.send('ack', self.zmq.NOBLOCK)
|
||||
@@ -119,9 +138,10 @@ class ComponentHost(Component):
|
||||
|
||||
class ParallelBuffer(Component):
|
||||
"""
|
||||
Connects to N PULL sockets, publishing all messages received to a PUB socket.
|
||||
Published messages are guaranteed to be in chronological order based on message property dt.
|
||||
Expects to be instantiated in one execution context (thread, process, etc) and run in another.
|
||||
Connects to N PULL sockets, publishing all messages received to a PUB
|
||||
socket. Published messages are guaranteed to be in chronological order
|
||||
based on message property dt. Expects to be instantiated in one execution
|
||||
context (thread, process, etc) and run in another.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
@@ -134,6 +154,7 @@ class ParallelBuffer(Component):
|
||||
self.ds_finished_counter = 0
|
||||
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return "FEED"
|
||||
|
||||
@@ -141,16 +162,16 @@ class ParallelBuffer(Component):
|
||||
self.data_buffer[source_id] = []
|
||||
|
||||
def open(self):
|
||||
self.pull_socket, self.poller = self.bind_data()
|
||||
self.feed_socket = self.bind_feed()
|
||||
self.pull_socket = self.bind_data()
|
||||
self.feed_socket = self.bind_feed()
|
||||
|
||||
def do_work(self):
|
||||
# wait for synchronization reply from the host
|
||||
socks = dict(self.poller.poll(2000)) #timeout after 2 seconds.
|
||||
socks = dict(self.poll.poll(self.heartbeat_timeout)) #timeout after 2 seconds.
|
||||
|
||||
if self.pull_socket in socks and socks[self.pull_socket] == self.zmq.POLLIN:
|
||||
message = self.pull_socket.recv()
|
||||
if self.is_done_message(message):
|
||||
if message == str(CONTROL_PROTOCOL.DONE):
|
||||
self.ds_finished_counter += 1
|
||||
if len(self.data_buffer) == self.ds_finished_counter:
|
||||
#drain any remaining messages in the buffer
|
||||
@@ -245,8 +266,8 @@ class MergedParallelBuffer(ParallelBuffer):
|
||||
ParallelBuffer.__init__(self)
|
||||
|
||||
def open(self):
|
||||
self.pull_socket, self.poller = self.bind_merge()
|
||||
self.feed_socket = self.bind_result()
|
||||
self.pull_socket = self.bind_merge()
|
||||
self.feed_socket = self.bind_result()
|
||||
|
||||
def next(self):
|
||||
"""Get the next merged message from the feed buffer."""
|
||||
@@ -263,6 +284,7 @@ class MergedParallelBuffer(ParallelBuffer):
|
||||
result[source] = cur['value']
|
||||
return result
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return "MERGE"
|
||||
|
||||
@@ -284,6 +306,7 @@ class BaseTransform(Component):
|
||||
self.state = {}
|
||||
self.state['name'] = name
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return self.state['name']
|
||||
|
||||
@@ -292,7 +315,7 @@ class BaseTransform(Component):
|
||||
Establishes zmq connections.
|
||||
"""
|
||||
#create the feed.
|
||||
self.feed_socket, self.poller = self.connect_feed()
|
||||
self.feed_socket = self.connect_feed()
|
||||
#create the result PUSH
|
||||
self.result_socket = self.connect_merge()
|
||||
|
||||
@@ -303,10 +326,10 @@ class BaseTransform(Component):
|
||||
- call transform (subclass' method) on event
|
||||
- send the transformed event
|
||||
"""
|
||||
socks = dict(self.poller.poll(2000)) #timeout after 2 seconds.
|
||||
socks = dict(self.poll.poll(self.heartbeat_timeout)) #timeout after 2 seconds.
|
||||
if self.feed_socket in socks and socks[self.feed_socket] == self.zmq.POLLIN:
|
||||
message = self.feed_socket.recv()
|
||||
if self.is_done_message(message):
|
||||
if message == str(CONTROL_PROTOCOL.DONE):
|
||||
self.signal_done()
|
||||
return
|
||||
|
||||
@@ -315,6 +338,7 @@ class BaseTransform(Component):
|
||||
#TODO: do we want to relay the datetime again? maybe drop this?
|
||||
#cur_state['dt'] = event['dt']
|
||||
cur_state['id'] = self.state['name']
|
||||
|
||||
self.result_socket.send(json.dumps(cur_state), self.zmq.NOBLOCK)
|
||||
|
||||
def transform(self, event):
|
||||
@@ -323,8 +347,9 @@ class BaseTransform(Component):
|
||||
|
||||
{name:"name of new transform", value: "value of new field"}
|
||||
|
||||
Transforms run in parallel and results are merged into a single map, so transform names must be unique.
|
||||
Best practice is to use the self.state object initialized from the transform configuration, and only set the
|
||||
Transforms run in parallel and results are merged into a single map, so
|
||||
transform names must be unique. Best practice is to use the self.state
|
||||
object initialized from the transform configuration, and only set the
|
||||
transformed value::
|
||||
|
||||
self.state['value'] = transformed_value
|
||||
@@ -352,6 +377,7 @@ class DataSource(Component):
|
||||
self.id = source_id
|
||||
self.cur_event = None
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return self.id
|
||||
|
||||
|
||||
+54
-65
@@ -10,6 +10,7 @@ class Controller(object):
|
||||
|
||||
def __init__(self, pull_socket, pub_socket, context=None, logging = None):
|
||||
|
||||
self.associated = []
|
||||
|
||||
if not context:
|
||||
self._ctx = zmq.Context()
|
||||
@@ -19,11 +20,6 @@ class Controller(object):
|
||||
self.pull_socket = pull_socket
|
||||
self.pub_socket = pub_socket
|
||||
|
||||
self.pull = self._ctx.socket(zmq.PULL)
|
||||
self.pub = self._ctx.socket(zmq.PUB)
|
||||
|
||||
self.associated = [self.pull, self.pub]
|
||||
|
||||
if logging:
|
||||
self.logging = logging
|
||||
self.dologging = True
|
||||
@@ -34,61 +30,71 @@ class Controller(object):
|
||||
self.success = 0
|
||||
self.failed = 0
|
||||
|
||||
try:
|
||||
self.pull.bind(pull_socket)
|
||||
except zmq.ZMQError:
|
||||
raise Exception('Cannot not bind on %s' % pull_socket)
|
||||
|
||||
try:
|
||||
self.pub.bind(pub_socket)
|
||||
except zmq.ZMQError:
|
||||
raise Exception('Cannot not bind on %s' % pub_socket)
|
||||
|
||||
def run(self, debug_step=False, stats=True):
|
||||
def run(self, debug=False):
|
||||
self.polling = True
|
||||
|
||||
if self.debug or debug_step:
|
||||
return self._poll_verbose(True, stats)
|
||||
else:
|
||||
return self._poll(False, stats)
|
||||
#if debug:
|
||||
return self._poll()
|
||||
#else:
|
||||
#return self._poll_fast()
|
||||
|
||||
def _poll_fast(self):
|
||||
"""
|
||||
C version of the polling forwarder.
|
||||
"""
|
||||
zmq.device(zmq.FORWARDER, self.pull, self.pub)
|
||||
|
||||
def _poll(self):
|
||||
"""
|
||||
Python version of the polling forwarder. With logging,
|
||||
mostly used for debugging.
|
||||
"""
|
||||
|
||||
self.pull = self._ctx.socket(zmq.PULL)
|
||||
self.pub = self._ctx.socket(zmq.PUB)
|
||||
|
||||
self.associated.extend([self.pull, self.pub])
|
||||
|
||||
self.pull.bind(self.pull_socket)
|
||||
self.pub.bind(self.pub_socket)
|
||||
|
||||
def _poll(self, debug_step, stats):
|
||||
while self.polling:
|
||||
try:
|
||||
self.logging.info('msg')
|
||||
self.pub.send(self.pull.recv())
|
||||
#self.pub.send(self.pull.recv(copy=False))
|
||||
except KeyboardInterrupt:
|
||||
self.polling = False
|
||||
break
|
||||
except zmq.ZMQError:
|
||||
self.polling = False
|
||||
break
|
||||
except Exception as e:
|
||||
# Its common to wrap these in wildcard exceptions so
|
||||
# that we don't loose messages, ever
|
||||
self.logging.error(str(e))
|
||||
if self.logging:
|
||||
self.logging.error(str(e))
|
||||
self.failed += 1
|
||||
continue
|
||||
|
||||
def _poll_verbose(self, debug_step, stats):
|
||||
while self.polling:
|
||||
try:
|
||||
if debug_step:
|
||||
msg = self.pull.recv(copy=False)
|
||||
if self.dologging:
|
||||
self.logging.info(msg)
|
||||
self.pub.send(msg)
|
||||
self.success += 1
|
||||
except KeyboardInterrupt:
|
||||
self.polling = False
|
||||
break
|
||||
except Exception as e:
|
||||
# Its common to wrap these in wildcard exceptions so
|
||||
# that we don't loose messages, ever
|
||||
self.logging.error(str(e))
|
||||
self.failed += 1
|
||||
continue
|
||||
def message_sender(self):
|
||||
"""
|
||||
Spin off a socket used for sending messages to this
|
||||
controller.
|
||||
"""
|
||||
s = self._ctx.socket(zmq.PUSH)
|
||||
s.connect(self.pull_socket)
|
||||
self.associated.append(s)
|
||||
return s
|
||||
|
||||
def qos(self):
|
||||
return float(self.success) / (self.success + self.failed)
|
||||
def message_listener(self):
|
||||
"""
|
||||
Spin off a socket used for receiving messages from this
|
||||
controller.
|
||||
"""
|
||||
s = self._ctx.socket(zmq.SUB)
|
||||
s.connect(self.pub_socket)
|
||||
s.setsockopt(zmq.SUBSCRIBE, '')
|
||||
self.associated.append(s)
|
||||
return s
|
||||
|
||||
def destroy(self):
|
||||
"""
|
||||
@@ -105,25 +111,8 @@ class Controller(object):
|
||||
def __del__(self):
|
||||
self.destroy()
|
||||
|
||||
def message_sender(self):
|
||||
"""
|
||||
Spin off a socket used for sending messages to this
|
||||
controller.
|
||||
"""
|
||||
s = self._ctx.socket(zmq.PUSH)
|
||||
s.connect(self.pull_socket)
|
||||
s.setsockopt(zmq.LINGER, -1)
|
||||
self.associated.append(s)
|
||||
return s
|
||||
|
||||
def message_listener(self):
|
||||
"""
|
||||
Spin off a socket used for receiving messages from this
|
||||
controller.
|
||||
"""
|
||||
s = self._ctx.socket(zmq.SUB)
|
||||
s.connect(self.pub_socket)
|
||||
s.setsockopt(zmq.SUBSCRIBE, '')
|
||||
self.associated.append(s)
|
||||
return s
|
||||
def qos(self):
|
||||
if not self.debug:
|
||||
return
|
||||
return float(self.success) / (self.success + self.failed)
|
||||
|
||||
|
||||
+145
-1
@@ -1,3 +1,147 @@
|
||||
#import msgpack
|
||||
"""
|
||||
The messaging protocol for Zipline.
|
||||
|
||||
Asserts are in place because any protocol error corresponds to a
|
||||
programmer error so we want it to fail fast and in an obvious way
|
||||
so it doesn't happen again. ZeroMQ follows the same philosophy.
|
||||
|
||||
Notes
|
||||
=====
|
||||
|
||||
Msgpack
|
||||
-------
|
||||
Msgpack is the fastest seriaization protocol in Python at the
|
||||
moment. Its 100% C is typically orders of magnitude faster than
|
||||
json and pickle making it awesome for ZeroMQ.
|
||||
|
||||
You can only serialize Python structural primitives: strings,
|
||||
numeric types, dicts, tuples and lists. Any any recursive
|
||||
combinations of these.
|
||||
|
||||
Basically every basestring in Python corresponds to valid
|
||||
msgpack message since the protocol is highly error tolerant.
|
||||
Just keep in mind that if you ever unpack a raw msgpack string
|
||||
make sure it looks like what you intend and/or catch ValueError
|
||||
and TypeError exceptions.
|
||||
|
||||
It also has the nice benefit of never invoking ``eval`` ( unlike
|
||||
json and pickle) which is a major security boon since it is
|
||||
impossible to arbitrary code for evaluation through messages.
|
||||
|
||||
UltraJSON
|
||||
---------
|
||||
For anything going to the browser UltraJSON is the fastest
|
||||
serializer, its mostly C as well.
|
||||
|
||||
The same domain of serialization as msgpack applies: Python
|
||||
structural primitives. It also has the additional constraint
|
||||
that anything outside of UTF8 can cause serious problems, so if
|
||||
you have a strong desire to JSON encode ancient Sanskrit
|
||||
( admit it, we all do ), just say no.
|
||||
|
||||
"""
|
||||
|
||||
import msgpack
|
||||
#import ujson
|
||||
#import ultrajson_numpy
|
||||
|
||||
from ctypes import Structure, c_ubyte
|
||||
|
||||
def Enum(*options):
|
||||
"""
|
||||
Fast enums are very important when we want really tight zmq
|
||||
loops. These are probably going to evolve into pure C structs
|
||||
anyways so might as well get going on that.
|
||||
"""
|
||||
class cstruct(Structure):
|
||||
_fields_ = [(o, c_ubyte) for o in options]
|
||||
return cstruct(*range(len(options)))
|
||||
|
||||
def FrameExceptionFactory(name):
|
||||
"""
|
||||
Exception factory with a closure around the frame class name.
|
||||
"""
|
||||
class InvalidFrame(Exception):
|
||||
def __init__(self, got):
|
||||
self.got = got
|
||||
def __str__(self):
|
||||
return "Invalid {framcls} Frame: {got}".format(
|
||||
framecls = name,
|
||||
got = self.got,
|
||||
)
|
||||
|
||||
class namedict(object):
|
||||
"""
|
||||
So that you can use:
|
||||
|
||||
foo.BAR
|
||||
-- or --
|
||||
foo['BAR']
|
||||
|
||||
For more complex strcuts use collections.namedtuple:
|
||||
"""
|
||||
|
||||
def __init__(self, dct):
|
||||
self.__dict__.update(dct)
|
||||
|
||||
# ================
|
||||
# Control Protocol
|
||||
# ================
|
||||
|
||||
INVALID_CONTROL_FRAME = FrameExceptionFactory('CONTROL')
|
||||
|
||||
CONTROL_PROTOCOL = Enum(
|
||||
'INIT' , # 0 - req
|
||||
'INFO' , # 1 - req
|
||||
'STATUS' , # 2 - req
|
||||
'SHUTDOWN' , # 3 - req
|
||||
'KILL' , # 4 - req
|
||||
|
||||
'OK' , # 5 - rep
|
||||
'DONE' , # 6 - rep
|
||||
'EXCEPTION' , # 7 - rep
|
||||
)
|
||||
|
||||
def CONTROL_FRAME(id, status):
|
||||
assert isinstance(basestring, id)
|
||||
assert isinstance(int, status)
|
||||
|
||||
return msgpack.dumps(tuple([id, status]))
|
||||
|
||||
def CONTORL_UNFRAME(msg):
|
||||
assert isinstance(basestring, msg)
|
||||
|
||||
try:
|
||||
id, status = msgpack.loads(msg)
|
||||
assert isinstance(basestring, id)
|
||||
assert isinstance(int, status)
|
||||
|
||||
return id, status
|
||||
except TypeError:
|
||||
raise INVALID_CONTROL_FRAME(msg)
|
||||
except ValueError:
|
||||
raise INVALID_CONTROL_FRAME(msg)
|
||||
#except AssertionError:
|
||||
#raise INVALID_CONTROL_FRAME(msg)
|
||||
|
||||
# ==================
|
||||
# Heartbeat Protocol
|
||||
# ==================
|
||||
|
||||
# These encode the msgpack equivelant of 1 and 2. The heartbeat
|
||||
# frame should only be 1 byte on the wire.
|
||||
|
||||
HEARTBEAT_PROTOCOL = namedict({
|
||||
'REQ' : b'\x01',
|
||||
'REP' : b'\x02',
|
||||
})
|
||||
|
||||
# ==================
|
||||
# Component State
|
||||
# ==================
|
||||
|
||||
COMPONENT_STATE = Enum(
|
||||
'OK' , # 0
|
||||
'DONE' , # 1
|
||||
'EXCEPTION' , # 2
|
||||
)
|
||||
|
||||
@@ -2,6 +2,8 @@ import json
|
||||
import zipline.util as qutil
|
||||
import zipline.messaging as qmsg
|
||||
|
||||
from zipline.protocol import CONTROL_PROTOCOL
|
||||
|
||||
class TestClient(qmsg.Component):
|
||||
|
||||
def __init__(self, utest, expected_msg_count=0):
|
||||
@@ -10,18 +12,22 @@ class TestClient(qmsg.Component):
|
||||
self.expected_msg_count = expected_msg_count
|
||||
self.utest = utest
|
||||
self.prev_dt = None
|
||||
self.heartbeat_timeout = 2000
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return "TEST_CLIENT"
|
||||
|
||||
def open(self):
|
||||
self.data_feed, self.poller = self.connect_result()
|
||||
self.data_feed = self.connect_result()
|
||||
|
||||
def do_work(self):
|
||||
socks = dict(self.poller.poll(2000)) #timeout after 2 seconds.
|
||||
socks = dict(self.poll.poll(self.heartbeat_timeout))
|
||||
|
||||
if self.data_feed in socks and socks[self.data_feed] == self.zmq.POLLIN:
|
||||
msg = self.data_feed.recv()
|
||||
if(self.is_done_message(msg)):
|
||||
|
||||
if msg == str(CONTROL_PROTOCOL.DONE):
|
||||
qutil.LOGGER.info("Client is DONE!")
|
||||
self.signal_done()
|
||||
self.utest.assertEqual(self.expected_msg_count, self.received_count,
|
||||
|
||||
+28
-15
@@ -7,32 +7,45 @@ import datetime
|
||||
import pytz
|
||||
import logging
|
||||
|
||||
|
||||
LOGGER = logging.getLogger('QSimLogger')
|
||||
|
||||
def configure_logging(loglevel=logging.DEBUG):
|
||||
"""
|
||||
Configures zipline.util.LOGGER to write a rotating file
|
||||
(10M per file, 5 files) to `` /var/log/zipline.log ``.
|
||||
"""
|
||||
LOGGER.setLevel(loglevel)
|
||||
handler = logging.handlers.RotatingFileHandler(
|
||||
"/var/log/zipline/{lfn}.log".format(lfn="zipline"),
|
||||
maxBytes=10*1024*1024, backupCount=5
|
||||
)
|
||||
handler.setFormatter(logging.Formatter(
|
||||
"%(asctime)s %(levelname)s %(filename)s %(funcName)s - %(message)s",
|
||||
"%Y-%m-%d %H:%M:%S %Z")
|
||||
)
|
||||
LOGGER.addHandler(handler)
|
||||
LOGGER.info("logging started...")
|
||||
|
||||
def parse_date(dt_str):
|
||||
"""parse strings according to the same format as generated by format_date"""
|
||||
"""
|
||||
Parse strings according to the same format as generated by
|
||||
format_date.
|
||||
"""
|
||||
if(dt_str == None):
|
||||
return None
|
||||
parts = dt_str.split(".")
|
||||
dt = datetime.datetime.strptime(parts[0], '%Y/%m/%d-%H:%M:%S').replace(microsecond=int(parts[1]+"000")).replace(tzinfo = pytz.utc)
|
||||
dt = datetime.datetime.strptime(parts[0], '%Y/%m/%d-%H:%M:%S').replace(
|
||||
microsecond=int(parts[1]+"000")).replace(tzinfo = pytz.utc
|
||||
)
|
||||
return dt
|
||||
|
||||
def format_date(dt):
|
||||
"""Format the date into a date with millesecond resolution and string/alphabetical
|
||||
sorting that is equivalent to datetime sorting"""
|
||||
"""
|
||||
Format the date into a date with millesecond resolution and
|
||||
string/alphabetical sorting that is equivalent to datetime sorting.
|
||||
"""
|
||||
if(dt == None):
|
||||
return None
|
||||
dt_str = dt.strftime('%Y/%m/%d-%H:%M:%S') + "." + str(dt.microsecond / 1000)
|
||||
return dt_str
|
||||
|
||||
def configure_logging(loglevel=logging.DEBUG):
|
||||
"""configures zipline.util.LOGGER to write a rotating file (10M per file, 5 files) to /var/log/zipline.log"""
|
||||
LOGGER.setLevel(loglevel)
|
||||
handler = logging.handlers.RotatingFileHandler(
|
||||
"/var/log/zipline/{lfn}.log".format(lfn="zipline"),
|
||||
maxBytes=10*1024*1024, backupCount=5)
|
||||
handler.setFormatter(logging.Formatter(
|
||||
"%(asctime)s %(levelname)s %(filename)s %(funcName)s - %(message)s","%Y-%m-%d %H:%M:%S %Z"))
|
||||
LOGGER.addHandler(handler)
|
||||
LOGGER.info("logging started...")
|
||||
Reference in New Issue
Block a user