Tweak tests.

This commit is contained in:
Stephen Diehl
2012-02-20 16:07:04 -05:00
parent 1cabd30158
commit 64383f865d
2 changed files with 42 additions and 47 deletions
+6 -10
View File
@@ -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))
+36 -37
View File
@@ -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)