mirror of
https://github.com/wassname/catalyst.git
synced 2026-08-12 11:50:11 +08:00
Tweak tests.
This commit is contained in:
+6
-10
@@ -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))
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user