Merge branch 'master' of github.com:quantopian/zipline

This commit is contained in:
fawce
2012-02-14 16:00:37 -05:00
5 changed files with 77 additions and 55 deletions
+57 -38
View File
@@ -2,7 +2,7 @@
Provides simulated data feed services...
"""
import multiprocessing
import zmq
from gevent_zeromq import zmq
import json
import copy
import threading
@@ -20,28 +20,34 @@ class SimulatorBase(object):
def __init__(self, sources, transforms, client, feed=None, merge=None):
"""
"""
self.sources = sources
self.transforms = transforms
self.client = client
self.merge = None
self.feed = None
self.context = None
self.sync_context = None
self.sync_socket = None
self.sync_register = {}
self.sync_address = "tcp://127.0.0.1:{port}".format(port=10100)
self.data_address = "tcp://127.0.0.1:{port}".format(port=10101)
self.feed_address = "tcp://127.0.0.1:{port}".format(port=10102)
self.merge_address = "tcp://127.0.0.1:{port}".format(port=10103)
self.result_address = "tcp://127.0.0.1:{port}".format(port=10104)
self.sources = sources
self.transforms = transforms
self.client = client
self.merge = None
self.feed = None
self.context = None
self.sync_context = None
self.sync_socket = None
self.sync_register = {}
self.sync_address = "tcp://127.0.0.1:{port}".format(port=10100)
self.data_address = "tcp://127.0.0.1:{port}".format(port=10101)
self.feed_address = "tcp://127.0.0.1:{port}".format(port=10102)
self.merge_address = "tcp://127.0.0.1:{port}".format(port=10103)
self.result_address = "tcp://127.0.0.1:{port}".format(port=10104)
self.timeout = datetime.timedelta(seconds=1)
self.performance_address = "tcp://127.0.0.1:{port}".format(port=10105)
self.timeout = datetime.timedelta(seconds=5)
#workaround for defect in threaded use of strptime: http://bugs.python.org/issue11108
qutil.parse_date("2012/02/13-10:04:28.114")
if(feed == None):
self.feed = DataFeed(self.sources.keys(), self.data_address, self.feed_address, qmsg.Sync(self,"DataFeed"))
self.feed = DataFeed(self.sources.keys(),
self.data_address,
self.feed_address,
self.performance_address,
qmsg.Sync(self,"DataFeed"))
else:
self.feed = feed
@@ -83,6 +89,7 @@ class SimulatorBase(object):
client_proc = self.launch_component("client", self.client)
qutil.LOGGER.info("client process launched")
qutil.LOGGER.info("sync register starting with {count} members: {reg}".format(count=len(self.sync_register), reg=self.sync_register))
self.sync_components()
#client_proc.join() #wait for client to complete processing
@@ -93,8 +100,9 @@ class SimulatorBase(object):
self.sync_register[sync_id] = datetime.datetime.utcnow()
def unregister_sync(self, sync_id):
qutil.LOGGER.info("unregistering {sync_id}".format(sync_id=sync_id))
del(self.sync_register[sync_id])
def is_timed_out(self):
cur_time = datetime.datetime.utcnow()
if(len(self.sync_register) == 0):
@@ -102,7 +110,7 @@ class SimulatorBase(object):
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}".format(source=source))
qutil.LOGGER.info("Time out for {source}. Current registery: {reg}".format(source=source, reg=self.sync_register))
return True
return False
@@ -112,7 +120,7 @@ class SimulatorBase(object):
qutil.LOGGER.info("waiting for all datasources and clients to be ready")
self.sync_socket = self.context.socket(zmq.REP)
self.sync_socket.bind(self.sync_address)
self.sync_socket.setsockopt(zmq.LINGER,0)
#self.sync_socket.setsockopt(zmq.LINGER,0)
self.poller = zmq.Poller()
self.poller.register(self.sync_socket, zmq.POLLIN)
@@ -132,9 +140,9 @@ class SimulatorBase(object):
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.sync_socket.send('ack', zmq.NOBLOCK)
except:
continue
qutil.LOGGER.exception("Exception in sync components loop")
self.sync_socket.close()
qutil.LOGGER.info("simulator heartbeat stopped.")
@@ -164,18 +172,20 @@ class ProcessSimulator(SimulatorBase):
class DataFeed(object):
def __init__(self, source_list, data_address, feed_address, sync):
def __init__(self, source_list, data_address, feed_address, performance_address, sync):
"""
:source_list: list of data source IDs
"""
self.feed_address = feed_address
self.data_address = data_address
self.data_buffer = qmsg.ParallelBuffer(source_list)
self.sync = sync
self.feed_socket = None
self.data_socket = None
self.context = None
self.poller = None
self.feed_address = feed_address
self.performance_address = performance_address
self.data_address = data_address
self.data_buffer = qmsg.ParallelBuffer(source_list)
self.sync = sync
self.feed_socket = None
self.data_socket = None
self.perf_socket = None
self.context = None
self.poller = None
def open(self):
@@ -189,18 +199,24 @@ class DataFeed(object):
#create the feed
self.feed_socket = self.context.socket(zmq.PUB)
self.feed_socket.bind(self.feed_address)
self.feed_socket.setsockopt(zmq.LINGER,0)
#self.feed_socket.setsockopt(zmq.LINGER,0)
self.data_buffer.out_socket = self.feed_socket
self.poller = zmq.Poller()
self.poller.register(self.data_socket, zmq.POLLIN)
#create the performance results push
self.perf_socket = self.context.socket(zmq.PUSH)
self.perf_socket.bind(self.performance_address)
#self.perf_socket.setsockopt(zmq.LINGER,0)
self.sync.open()
def close(self):
try:
self.data_socket.close()
self.feed_socket.close()
self.perf_socket.close()
self.sync.close()
except:
qutil.LOGGER.exception("Error closing DataFeed")
@@ -225,12 +241,14 @@ class DataFeed(object):
else:
self.data_buffer.append(event[u's'], event)
self.data_buffer.send_next()
#self.perf_socket.send(message, zmq.NOBLOCK)
#drain any remaining messages in the buffer
self.data_buffer.drain()
#send the DONE message
self.feed_socket.send("DONE")
self.feed_socket.send("DONE", zmq.NOBLOCK)
qutil.LOGGER.info("received {n} messages, sent {m} messages".format(n=self.data_buffer.received_count,
m=self.data_buffer.sent_count))
@@ -300,7 +318,7 @@ class BaseTransform(object):
#create the result PUSH
self.result_socket = self.context.socket(zmq.PUSH)
self.result_socket.connect(self.merge_address)
self.result_socket.setsockopt(zmq.LINGER,0)
#self.result_socket.setsockopt(zmq.LINGER,0)
self.sync.open()
@@ -319,14 +337,14 @@ class BaseTransform(object):
message = self.feed_socket.recv()
if(message == "DONE"):
qutil.LOGGER.info("{name} received the Done message from the feed".format(name=self.state['name']))
self.result_socket.send("DONE")
self.result_socket.send("DONE", zmq.NOBLOCK)
break
self.received_count += 1
event = json.loads(message)
cur_state = self.transform(event)
cur_state['dt'] = event['dt']
cur_state['name'] = self.state['name']
self.result_socket.send(json.dumps(cur_state))
self.result_socket.send(json.dumps(cur_state), zmq.NOBLOCK)
self.sent_count += 1
def close(self):
@@ -405,7 +423,7 @@ class TransformsMerge(object):
#create the result PUSH
self.result_socket = self.context.socket(zmq.PUSH)
self.result_socket.bind(self.result_address)
self.result_socket.setsockopt(zmq.LINGER,0)
#self.result_socket.setsockopt(zmq.LINGER,0)
#create the transform PULL.
self.transform_socket = self.context.socket(zmq.PULL)
@@ -424,6 +442,7 @@ class TransformsMerge(object):
Close all zmq sockets and context.
"""
try:
self.sync.close()
self.transform_socket.close()
self.feed_socket.close()
self.result_socket.close()
@@ -472,7 +491,7 @@ class TransformsMerge(object):
self.data_buffer.drain()
#signal to client that we're done
self.result_socket.send("DONE")
self.result_socket.send("DONE", zmq.NOBLOCK)
+5 -4
View File
@@ -3,7 +3,7 @@ Commonly used messaging components.
"""
import json
import uuid
import zmq
from gevent_zeromq import zmq
import zipline.util as qutil
@@ -72,7 +72,7 @@ class ParallelBuffer(object):
event = self.next()
if(event != None):
self.out_socket.send(json.dumps(event))
self.out_socket.send(json.dumps(event), zmq.NOBLOCK)
self.sent_count += 1
@@ -130,12 +130,13 @@ class Sync(object):
"""Confirm readiness with the Host."""
try:
# send a synchronization request to the host
self.sync_socket.send(self.sync_id + ":RUNNING", zmq.NOBLOCK)
self.sync_socket.send(self.sync_id + ":RUNNING")
# wait for synchronization reply from the host
socks = dict(self.poller.poll(2000)) #timeout after 2 seconds.
if self.sync_socket in socks and socks[self.sync_socket] == zmq.POLLIN:
message = self.sync_socket.recv()
return True
except:
qutil.LOGGER.exception("exception in confirmation for {source}. Exiting.".format(source=self.sync_id))
@@ -143,7 +144,7 @@ class Sync(object):
def close(self):
try:
self.sync_socket.send(self.sync_id + ":DONE", zmq.NOBLOCK)
self.sync_socket.send(self.sync_id + ":DONE")
self.sync_socket.close()
except:
qutil.LOGGER.exception("Error closing Sync object")
+4 -3
View File
@@ -2,7 +2,7 @@
Provides data handlers that can push messages to a zipline.core.DataFeed
"""
import datetime
import zmq
from gevent_zeromq import zmq
import json
import random
@@ -33,7 +33,7 @@ class DataSource(object):
#create the data sink. Based on http://zguide.zeromq.org/py:tasksink2
self.data_socket = self.context.socket(zmq.PUSH)
self.data_socket.connect(self.data_address)
self.data_socket.setsockopt(zmq.LINGER,0)
#self.data_socket.setsockopt(zmq.LINGER,0)
self.sync.open()
@@ -57,7 +57,6 @@ class DataSource(object):
sets source_id and type properties in the dict
sends to the data_socket.
"""
self.sync.confirm()
event['s'] = self.source_id
event['type'] = 'event'
self.data_socket.send(json.dumps(event), zmq.NOBLOCK)
@@ -98,6 +97,8 @@ class RandomEquityTrades(DataSource):
price = random.uniform(5.0, 50.0)
for i in range(self.count):
if not self.sync.confirm():
break
price = price + random.uniform(-0.05, 0.05)
event = {'sid':self.sid,
'dt':qutil.format_date(trade_start + (minute * i)),
+2 -2
View File
@@ -1,4 +1,4 @@
import zmq
from gevent_zeromq import zmq
import json
import zipline.util as qutil
@@ -55,7 +55,7 @@ class TestClient(object):
self.sync.close()
except:
self.error = True
qutil.LOGGER.exception("**********************Error in test client.")
qutil.LOGGER.exception("Error in test client.")
finally:
self.context.destroy()
+9 -8
View File
@@ -15,18 +15,19 @@ import zipline.messaging as qmsg
from zipline.test.client import TestClient
qutil.configure_logging()
class MessagingTestCase(unittest.TestCase):
"""Tests the message passing: datasources -> feed -> transforms -> merge -> client"""
def setUp(self):
"""generate some config objects for the datafeed, sources, and transforms."""
qutil.configure_logging()
pass
def get_simulator(self, sources, transforms, client, feed=None, merge=None):
return ProcessSimulator(sources, transforms, client, feed=feed, merge=merge)
def test_sources_only(self):
def dtest_sources_only(self):
"""streams events from two data sources, no transforms."""
ret1 = RandomEquityTrades(133, "ret1", 400)
@@ -47,13 +48,13 @@ class MessagingTestCase(unittest.TestCase):
verify message count at client.
"""
ret1 = RandomEquityTrades(133, "ret1", 400)
ret2 = RandomEquityTrades(134, "ret2", 400)
ret1 = RandomEquityTrades(133, "ret1", 5000)
ret2 = RandomEquityTrades(134, "ret2", 5000)
sources = {"ret1":ret1, "ret2":ret2}
mavg1 = MovingAverage("mavg1", 30)
mavg2 = MovingAverage("mavg2", 60)
transforms = {"mavg1":mavg1, "mavg2":mavg2}
client = TestClient(self, expected_msg_count=800)
client = TestClient(self, expected_msg_count=10000)
sim = self.get_simulator(sources, transforms, client)
sim.simulate()
@@ -69,14 +70,14 @@ class MessagingTestCase(unittest.TestCase):
transforms = {"mavg1":mavg1, "mavg2":mavg2}
client = TestClient(self, expected_msg_count=0)
sim = self.get_simulator(sources, transforms, client)
sim.feed = DataFeedErr(sources.keys(), sim.data_address, sim.feed_address, qmsg.Sync(sim, "DataFeedErrorGenerator"))
sim.feed = DataFeedErr(sources.keys(), sim.data_address, sim.feed_address, sim.performance_address, qmsg.Sync(sim, "DataFeedErrorGenerator"))
sim.simulate()
class DataFeedErr(DataFeed):
"""Helper class for testing, simulates exceptions inside the DataFeed"""
def __init__(self, source_list, data_address, feed_address, sync):
DataFeed.__init__(self, source_list, data_address, feed_address, sync)
def __init__(self, source_list, data_address, feed_address, perf_address, sync):
DataFeed.__init__(self, source_list, data_address, feed_address, perf_address, sync)
def handle_all(self):
#time.sleep(1000)