diff --git a/zipline/test/client.py b/zipline/test/client.py index 0fad7187..3b975ec2 100644 --- a/zipline/test/client.py +++ b/zipline/test/client.py @@ -3,20 +3,20 @@ import zipline.util as qutil import zipline.messaging as qmsg class TestClient(qmsg.Component): - + def __init__(self, utest, expected_msg_count=0): qmsg.Component.__init__(self) self.received_count = 0 self.expected_msg_count = expected_msg_count self.utest = utest self.prev_dt = None - + def get_id(self): - return "TEST_CLIENT" - + return "TEST_CLIENT" + def open(self): self.data_feed, self.poller = self.connect_result() - + def do_work(self): socks = dict(self.poller.poll(2000)) #timeout after 2 seconds. if self.data_feed in socks and socks[self.data_feed] == self.zmq.POLLIN: @@ -28,7 +28,7 @@ class TestClient(qmsg.Component): "The client should have received ({n}) the same number of messages as the feed sent ({m})." .format(n=self.received_count, m=self.expected_msg_count)) return - + self.received_count += 1 event = json.loads(msg) if(self.prev_dt != None): @@ -38,7 +38,3 @@ class TestClient(qmsg.Component): self.prev_dt = event['dt'] if(self.received_count % 100 == 0): qutil.LOGGER.info("received {n} messages".format(n=self.received_count)) - - - - \ No newline at end of file diff --git a/zipline/test/test_messaging.py b/zipline/test/test_messaging.py index 15e91396..1a42faa1 100644 --- a/zipline/test/test_messaging.py +++ b/zipline/test/test_messaging.py @@ -3,57 +3,56 @@ Test suite for the messaging infrastructure of QSim. """ #don't worry about excessive public methods pylint: disable=R0904 -import unittest2 as unittest -import multiprocessing -import time +import zipline.util as qutil +import zipline.messaging as qmsg from zipline.simulator import ThreadSimulator, ProcessSimulator from zipline.transforms.technical import MovingAverage from zipline.sources import RandomEquityTrades -import zipline.util as qutil -import zipline.messaging as qmsg from zipline.test.client import TestClient -qutil.configure_logging() +# Should not inherit form TestCase since test runners will pick +# it up as a test. -class SimulatorTestCase(unittest.TestCase): - """Tests the message passing: datasources -> feed -> transforms -> merge -> client""" +class SimulatorTestCase(object): def setUp(self): - """generate some config objects for the datafeed, sources, and transforms.""" - self.addresses = {'sync_address' : "tcp://127.0.0.1:10100", - 'data_address' : "tcp://127.0.0.1:10101", - 'feed_address' : "tcp://127.0.0.1:10102", - 'merge_address' : "tcp://127.0.0.1:10103", - 'result_address' : "tcp://127.0.0.1:10104" - } - - self.addressesblarg = "test" - - def get_simulator(self): - return ThreadSimulator(self.addresses) + qutil.configure_logging() + """ + Generate some config objects for the datafeed, sources, and transforms. + """ + + self.addresses = { + 'sync_address' : "tcp://127.0.0.1:10100", + 'data_address' : "tcp://127.0.0.1:10101", + 'feed_address' : "tcp://127.0.0.1:10102", + 'merge_address' : "tcp://127.0.0.1:10103", + 'result_address' : "tcp://127.0.0.1:10104" + } + + self.addressesblarg = "test" + + def get_simulator(self): + raise NotImplementedError + #return ThreadSimulator(self.addresses) def test_sources_only(self): - """streams events from two data sources, no transforms.""" + sim = self.get_simulator() ret1 = RandomEquityTrades(133, "ret1", 400) ret2 = RandomEquityTrades(134, "ret2", 400) client = TestClient(self, expected_msg_count=800) sim.register_components([ret1, ret2, client]) sim.simulate() - - self.assertEqual(sim.feed.pending_messages(), 0, + + self.assertEqual(sim.feed.pending_messages(), 0, "The feed should be drained of all messages, found {n} remaining." .format(n=sim.feed.pending_messages())) - - + + def test_transforms(self): - """ - 2 datasources -> feed -> 2 moving average transforms -> transform merge -> testclient - verify message count at client. - """ sim = self.get_simulator() ret1 = RandomEquityTrades(133, "ret1", 5000) ret2 = RandomEquityTrades(134, "ret2", 5000) @@ -62,9 +61,9 @@ class SimulatorTestCase(unittest.TestCase): client = TestClient(self, expected_msg_count=10000) sim.register_components([ret1, ret2, mavg1, mavg2, client]) sim.simulate() - + self.assertEqual(sim.feed.pending_messages(), 0, "The feed should be drained of all messages.") - + def dtest_error_in_feed(self): ret1 = RandomEquityTrades(133, "ret1", 400) ret2 = RandomEquityTrades(134, "ret2", 400) @@ -76,10 +75,10 @@ class SimulatorTestCase(unittest.TestCase): sim = self.get_simulator(sources, transforms, client) sim.feed = DataFeedErr(sources.keys(), sim.data_address, sim.feed_address, sim.performance_address, qmsg.Sync(sim, "DataFeedErrorGenerator")) sim.simulate() - - -class ProcessSimulatorTestCase(SimulatorTestCase): - - def get_simulator(self): - return ProcessSimulator(self.addresses) + + +#class ProcessSimulatorTestCase(SimulatorTestCase): + + #def get_simulator(self): + #return ProcessSimulator(self.addresses)