From 536a1e7fdc9f655883921da3e3ff63f902e88176 Mon Sep 17 00:00:00 2001 From: fawce Date: Wed, 29 Feb 2012 23:38:46 -0500 Subject: [PATCH] protocol tests passing, orders test almost passing, but looking like we need to extend the merge component to add transaction and order specific logic. --- zipline/finance/trading.py | 104 +++++++++++++++------------ zipline/messaging.py | 5 +- zipline/protocol.py | 131 +++++++++++++++++++++++++---------- zipline/test/client.py | 16 ++--- zipline/test/factory.py | 60 +--------------- zipline/test/test_finance.py | 53 ++++++++++++-- 6 files changed, 211 insertions(+), 158 deletions(-) diff --git a/zipline/finance/trading.py b/zipline/finance/trading.py index e3a97e53..a6b81d2e 100644 --- a/zipline/finance/trading.py +++ b/zipline/finance/trading.py @@ -17,7 +17,7 @@ class TradeSimulationClient(qmsg.Component): @property def get_id(self): - return "TRADING_CLIENT" + return str(zp.FINANCE_COMPONENT.TRADING_CLIENT) def open(self): self.result_feed = self.connect_result() @@ -25,21 +25,18 @@ class TradeSimulationClient(qmsg.Component): def do_work(self): #next feed event - (rlist, wlist, xlist) = select([self.result_feed], - [], - [self.result_feed], - timeout=self.heartbeat_timeout/100) #select timeout is in sec, use 10x - # - #no more orders, should be an error condition - if len(rlist) == 0 or len(xlist) > 0: - raise Exception("unexpected end of feed stream") - message = rlist[0].recv() - if message == str(zp.CONTROL_PROTOCOL.DONE): - self.signal_done() - return #leave open orders hanging? client requests for orders? + socks = dict(self.poll.poll(self.heartbeat_timeout)) + + if self.result_feed in socks and socks[self.result_feed] == self.zmq.POLLIN: + msg = self.result_feed.recv() + + if msg == str(zp.CONTROL_PROTOCOL.DONE): + qutil.LOGGER.info("Client is DONE!") + self.signal_done() + return - event = zp.MERGE_UNFRAME(message) - self._handle_event(event) + event = zp.MERGE_UNFRAME(message) + self._handle_event(event) def connect_order(self): return self.connect_push_socket(self.addresses['order_address']) @@ -71,12 +68,13 @@ class OrderDataSource(qmsg.DataSource): 'volume' : integer for volume } """ - zm.DataSource.__init__(self, str(zp.FINANCE_PROTOCOL.ORDER)) + qmsg.DataSource.__init__(self, zp.FINANCE_COMPONENT.ORDER_SOURCE) self.simulation_dt = simulation_dt self.last_iteration_duration = datetime.timedelta(seconds=0) + self.sent_count = 0 def get_type(self): - return str(zp.FINANCE_PROTOCOL.ORDER) + return zp.FINANCE_COMPONENT.ORDER_SOURCE def open(self): qmsg.DataSource.open(self) @@ -85,16 +83,19 @@ class OrderDataSource(qmsg.DataSource): def bind_order(self): return self.bind_pull_socket(self.addresses['order_address']) - def do_work(self): + def do_work(self): #mark the start time for client's processing of this event. self.event_start = datetime.datetime.utcnow() - self.result_socket.send(zp.TRANSFORM_FRAME(str(zp.FINANCE_PROTOCOL.ORDER), self.simulation_dt), self.zmq.NOBLOCK) - self.simulation_dt = self.simulation_dt + self.last_iteration_duration + #TODO: if this is the first iteration, break deadlock by sending a dummy order + if(self.sent_count == 0): + self.send_dummy() + #pull all orders from client. orders = [] order_dt = None + count = 0 while True: (rlist, wlist, xlist) = select([self.order_socket], [], @@ -115,21 +116,35 @@ class OrderDataSource(qmsg.DataSource): self.last_iteration_duration = datetime.datetime.utcnow() - self.event_start dt = self.simulation_dt + self.last_iteration_duration - order_event = zp.namedict({"sid":sid, "amount":amount, "dt":dt, source_id=self.get_id}) + order_event = zp.namedict({"sid":sid, "amount":amount, "dt":dt, "source_id":self.get_id, "type":zp.DATASOURCE_TYPE.ORDER}) - message = zp.DATASOURCE_FRAME(event) - self.data_socket.send(message) + self.send(order_event) + count += 1 + self.sent_count += 1 + #TODO: we have to send at least one dummy order per do_work iteration or the feed will block waiting for our messages. + if(count == 0): + self.send_dummy() + self.sent_count += 1 + + def send(self, order_event): + message = zp.DATASOURCE_FRAME(order_event) + self.data_socket.send(message) + + def send_dummy(self): + dt = self.simulation_dt + self.last_iteration_duration + dummy_order = zp.namedict({"sid":0, "amount":0, "dt":dt, "source_id":self.get_id, "type":zp.DATASOURCE_TYPE.ORDER}) + self.send(dummy_order) class TransactionSimulator(qmsg.BaseTransform): def __init__(self): - qmsg.BaseTransform.__init__(self, "TRANSACTION_SIM") + qmsg.BaseTransform.__init__(self, zp.TRANSFORM_TYPE.TRANSACTION) self.open_orders = {} self.order_count = 0 - self.tradeWindow = datetime.timedelta(seconds=30) + self.trade_windwo = datetime.timedelta(seconds=30) self.orderTTL = datetime.timedelta(days=1) self.volume_share = 0.05 self.commission = 0.03 @@ -139,12 +154,14 @@ class TransactionSimulator(qmsg.BaseTransform): Pulls one message from the event feed, then loops on orders until client sends DONE message. """ - if(event.type == zp.FINANCE_PROTOCOL.ORDER): + #TODO: need a way to send a placeholder txn, to avoid blocking merge... maybe customize merge to not block on txn? + if(event.type == zp.DATASOURCE_TYPE.ORDER): self.add_open_order(event.sid, event.amount) - self.state['value'] = self.average - elif(event.type == zp.FINANCE_PROTOCOL.TRADE): + self.state['value'] = self.create_transaction(0, 0, 0.0, event.dt, 1) + elif(event.type == zp.DATASOURCE_TYPE.TRADE): txn = apply_trade_to_open_orders(event) self.state['value'] = txn + #TODO: what to do if we get another kind of datasource event.type? return self.state @@ -155,17 +172,16 @@ class TransactionSimulator(qmsg.BaseTransform): """ amount = int(amount) if amount == 0: - qutil.LOGGER.debug("{title}:{id} requested to trade zero shares of {sid}".format(sid=sid, - title=self.hostedAlgo.algo.title, - id=self.hostedAlgo.algo.id)) + qutil.LOGGER.debug("requested to trade zero shares of {sid}".format(sid=sid)) return self.order_count += 1 order = zp.namedict({'sid' : sid, 'amount' : amount, - 'dt' : self.algo_time}, + 'dt' : self.algo_time, 'filled': 0, - 'direction': math.fabs(amount) / amount) + 'direction': math.fabs(amount) / amount + }) if(not self.open_orders.has_key(sid)): self.open_orders[sid] = [] @@ -175,7 +191,7 @@ class TransactionSimulator(qmsg.BaseTransform): if(event.volume == 0): #there are zero volume events bc some stocks trade less frequently than once per minute. - continue + return if self.open_orders.has_key(event.sid): orders = self.open_orders[event.sid] remaining_orders = [] @@ -184,7 +200,7 @@ class TransactionSimulator(qmsg.BaseTransform): for order in orders: #we're using minute bars, so allow orders within 30 seconds of the trade - if((order.dt - event.dt) < self.tradeWindow): + if((order.dt - event.dt) < self.trade_windwo): total_order += order.amount if(order.dt > dt): dt = order.dt @@ -206,16 +222,14 @@ class TransactionSimulator(qmsg.BaseTransform): def create_transaction(self, sid, amount, price, dt, direction): - if(amount != 0): - txn = {'sid' : sid, - 'amount' : amount, - 'dt' : dt, - 'price' : price, - 'back_test_run_id' : self.btRun.id, - 'transaction_cost' : -1*(price * amount), - 'commision' : self.commission * amount * direction} - - return namedict(txn) + txn = {'sid' : sid, + 'amount' : amount, + 'dt' : dt, + 'price' : price, + 'commission' : self.commission * amount * direction, + 'source_id' : zp.FINANCE_COMPONENT.TRANSACTION_SIM + } + return zp.namedict(txn) diff --git a/zipline/messaging.py b/zipline/messaging.py index 70268774..7e8565e6 100644 --- a/zipline/messaging.py +++ b/zipline/messaging.py @@ -409,8 +409,9 @@ class PassthroughTransform(BaseTransform): if message == str(CONTROL_PROTOCOL.DONE): self.signal_done() return - #message is already FEED_FRAMEd, send it as the value. - self.result_socket.send(zp.TRANSFORM_FRAME("PASSTHROUGH", message), self.zmq.NOBLOCK) + + #message is already FEED_FRAMEd, send it as the value. + self.result_socket.send(zp.TRANSFORM_FRAME("PASSTHROUGH", message), self.zmq.NOBLOCK) class DataSource(Component): diff --git a/zipline/protocol.py b/zipline/protocol.py index 22ca9712..fe074840 100644 --- a/zipline/protocol.py +++ b/zipline/protocol.py @@ -190,9 +190,11 @@ def DATASOURCE_FRAME(event): ::payload:: a msgpack string carrying the payload for the frame """ assert isinstance(event.source_id, basestring) - assert isinstance(event.type, basestring) - if(event.type == "TRADE"): + assert isinstance(event.type, int) + if(event.type == DATASOURCE_TYPE.TRADE): return msgpack.dumps(tuple([event.type, TRADE_FRAME(event)])) + elif(event.type == DATASOURCE_TYPE.ORDER): + return msgpack.dumps(tuple([event.type, ORDER_SOURCE_FRAME(event)])) else: raise INVALID_DATASOURCE_FRAME(str(event)) @@ -211,9 +213,11 @@ def DATASOURCE_UNFRAME(msg): """ try: ds_type, payload = msgpack.loads(msg) - assert isinstance(ds_type, basestring) - if(ds_type == "TRADE"): + assert isinstance(ds_type, int) + if(ds_type == DATASOURCE_TYPE.TRADE): return TRADE_UNFRAME(payload) + elif(ds_type == DATASOURCE_TYPE.ORDER): + return ORDER_SOURCE_UNFRAME(payload) else: raise INVALID_DATASOURCE_FRAME(msg) @@ -266,10 +270,8 @@ def TRANSFORM_FRAME(name, value): """ assert isinstance(name, basestring) assert value != None - - if(name == 'SIM_DT'): - value = PACK_ALGO_DT(value) - + if(name == TRANSFORM_TYPE.TRANSACTION): + value = TRANSACTION_FRAME(value) return msgpack.dumps(tuple([name, value])) def TRANSFORM_UNFRAME(msg): @@ -280,26 +282,16 @@ def TRANSFORM_UNFRAME(msg): name, value = msgpack.loads(msg) #TODO: anything we can do to assert more about the content of the dict? assert isinstance(name, basestring) - if(name == "PASSTHROUGH"): + if(name == TRANSFORM_TYPE.PASSTHROUGH): value = FEED_UNFRAME(value) - elif(name == "SIM_DT"): - value = UNPACK_ALGO_DT(value) + elif(name == TRANSFORM_TYPE.TRANSACTION): + value = TRANSACTION_UNFRAME(value) return namedict({name : value}) except TypeError: raise INVALID_TRANSFORM_FRAME(msg) except ValueError: raise INVALID_TRANSFORM_FRAME(msg) -def PACK_ALGO_DT(value): - value = namedict({'dt' : value}) - PACK_DATE(value) - return value.__dict__ - -def UNPACK_ALGO_DT(value): - value = namedict(value) - UNPACK_DATE(value) - return value.dt - # ================== # Merge Protocol # ================== @@ -312,10 +304,7 @@ def MERGE_FRAME(event): - type """ assert isinstance(event, namedict) - assert isinstance(event.dt, datetime.datetime) PACK_DATE(event) - if(event.has_attr('SIM_DT')): - event.SIM_DT = PACK_ALGO_DT(event.SIM_DT) payload = event.__dict__ return msgpack.dumps(payload) @@ -325,10 +314,6 @@ def MERGE_UNFRAME(msg): #TODO: anything we can do to assert more about the content of the dict? assert isinstance(payload, dict) payload = namedict(payload) - if(payload.has_attr('SIM_DT')): - payload.SIM_DT = UNPACK_ALGO_DT(payload.SIM_DT) - assert isinstance(payload.epoch, numbers.Integral) - assert isinstance(payload.micros, numbers.Integral) UNPACK_DATE(payload) return payload except TypeError: @@ -344,7 +329,7 @@ INVALID_ORDER_FRAME = FrameExceptionFactory('ORDER') INVALID_TRADE_FRAME = FrameExceptionFactory('TRADE') # ================== -# Trades +# Trades - Should only be called from inside DATASOURCE_ (UN)FRAME. # ================== def TRADE_FRAME(event): @@ -358,7 +343,7 @@ def TRADE_FRAME(event): """ assert isinstance(event, namedict) assert isinstance(event.source_id, basestring) - assert event.type == "TRADE" + assert event.type == DATASOURCE_TYPE.TRADE assert isinstance(event.sid, int) assert isinstance(event.price, float) assert isinstance(event.volume, int) @@ -372,8 +357,6 @@ def TRADE_UNFRAME(msg): assert isinstance(sid, int) assert isinstance(price, float) assert isinstance(volume, int) - assert isinstance(epoch, numbers.Integral) - assert isinstance(micros, numbers.Integral) rval = namedict({'sid' : sid, 'price' : price, 'volume' : volume, 'epoch' : epoch, 'micros' : micros, 'type' : source_type, 'source_id' : source_id}) UNPACK_DATE(rval) return rval @@ -383,7 +366,7 @@ def TRADE_UNFRAME(msg): raise INVALID_TRADE_FRAME(msg) # ========= -# Orders +# Orders - from client to order source # ========= def ORDER_FRAME(sid, amount): @@ -404,6 +387,66 @@ def ORDER_UNFRAME(msg): except ValueError: raise INVALID_ORDER_FRAME(msg) +# +# ================== +# TRANSACTIONS - Should only be called from inside TRANSFORM_(UN)FRAME. +# ================== + +def TRANSACTION_FRAME(event): + assert isinstance(event, namedict) + assert isinstance(event.sid, int) + assert isinstance(event.price, float) + assert isinstance(event.commission, float) + assert isinstance(event.amount, int) + PACK_DATE(event) + return msgpack.dumps(tuple([event.sid, event.price, event.amount, event.commission, event.epoch, event.micros])) + +def TRANSACTION_UNFRAME(msg): + try: + sid, price, amount, commission, epoch, micros = msgpack.loads(msg) + + assert isinstance(sid, int) + assert isinstance(price, float) + assert isinstance(commission, float) + assert isinstance(amount, int) + rval = namedict({'sid' : sid, 'price' : price, 'amount' : amount, 'commission':commission, 'epoch' : epoch, 'micros' : micros}) + UNPACK_DATE(rval) + return rval + except TypeError: + raise INVALID_TRADE_FRAME(msg) + except ValueError: + raise INVALID_TRADE_FRAME(msg) + + +# ========= +# Orders - from order source to feed +# - should only be called from inside DATASOURCE_(UN)FRAME +# ========= + +def ORDER_SOURCE_FRAME(event): + assert isinstance(event.sid, int) + assert isinstance(event.amount, int) #no partial shares... + assert isinstance(event.source_id, basestring) + assert event.type == DATASOURCE_TYPE.ORDER + PACK_DATE(event) + return msgpack.dumps(tuple([event.sid, event.amount, event.epoch, event.micros, event.source_id, event.type])) + + +def ORDER_SOURCE_UNFRAME(msg): + try: + sid, amount, epoch, micros, source_id, source_type = msgpack.loads(msg) + event = namedict({"sid":sid, "amount":amount, "epoch":epoch, "micros":micros, "source_id":source_id, "type":source_type}) + assert isinstance(sid, int) + assert isinstance(amount, int) + assert isinstance(source_id, basestring) + assert isinstance(source_type, int) + UNPACK_DATE(event) + return event + except TypeError: + raise INVALID_ORDER_FRAME(msg) + except ValueError: + raise INVALID_ORDER_FRAME(msg) + # ================= # Date Helpers # ================= @@ -428,8 +471,20 @@ def UNPACK_DATE(payload): return payload -FINANCE_PROTOCOL = Enum( - 'ORDER' , # 0 - 'TRANSACTION' , # 1 - 'TRADE' , # 2 - ) +DATASOURCE_TYPE = Enum( + 'ORDER' , + 'TRADE' , +) + +TRANSFORM_TYPE = namedict({ + 'TRANSACTION':'TRANSACTION', #needed? + 'PASSTHROUGH':'PASSTHROUGH' + }) + + +FINANCE_COMPONENT = namedict({ + 'TRADING_CLIENT':'TRADING_CLIENT', + 'PORTFOLIO_CLIENT':'PORTFOLIO_CLIENT', + 'ORDER_SOURCE':'ORDER_SOURCE', + 'TRANSACTION_SIM':'TRANSACTION_SIM' + }) diff --git a/zipline/test/client.py b/zipline/test/client.py index e9649977..02a8e4c1 100644 --- a/zipline/test/client.py +++ b/zipline/test/client.py @@ -6,12 +6,11 @@ from zipline.finance.trading import TradeSimulationClient import zipline.protocol as zp class TestClient(qmsg.Component): - """no-op client - Just connects to the merge and counts messages. compares received message count to the expected count.""" + """no-op client - Just connects to the merge and counts messages.""" - def __init__(self, expected_msg_count=0): + def __init__(self): qmsg.Component.__init__(self) self.received_count = 0 - self.expected_msg_count = expected_msg_count self.prev_dt = None self.heartbeat_timeout = 2000 @@ -31,9 +30,6 @@ class TestClient(qmsg.Component): if msg == str(zp.CONTROL_PROTOCOL.DONE): qutil.LOGGER.info("Client is DONE!") self.signal_done() - if(self.expected_msg_count > 0): - assert self.received_count == self.expected_msg_count - return self.received_count += 1 @@ -48,14 +44,16 @@ class TestClient(qmsg.Component): class TestTradingClient(TradeSimulationClient): - def __init__(self, count): + def __init__(self, sid, amount, order_count): TradeSimulationClient.__init__(self) - self.count = count + self.count = order_count + self.sid = sid + self.amount = amount self.incr = 0 def handle_events(self, event_queue): #place an order for 100 shares of sid:133 if(self.incr < self.count): - self.order(133, 100) + self.order(self.sid, self.amount) self.incr += 1 diff --git a/zipline/test/factory.py b/zipline/test/factory.py index a30979e5..2a460506 100644 --- a/zipline/test/factory.py +++ b/zipline/test/factory.py @@ -2,67 +2,13 @@ import datetime import pytz import zipline.util as qutil import zipline.finance.risk as risk +import zipline.protocol as zp -def createReturns(daycount, start): - i = 0 - test_range = [] - current = start.replace(tzinfo=pytz.utc) - one_day = datetime.timedelta(days = 1) - while i < daycount: - i += 1 - r = daily_return(current, random.random()) - test_range.append(r) - current = current + one_day - return [ x for x in test_range if(risk.trading_calendar.is_trading_day(x.date)) ] - -def createReturnsFromRange(start, end): - current = start.replace(tzinfo=pytz.utc) - end = end.replace(tzinfo=pytz.utc) - one_day = datetime.timedelta(days = 1) - test_range = [] - i = 0 - while current <= end: - current = current + one_day - if(not risk.trading_calendar.is_trading_day(current)): - continue - r = daily_return(current, random.random()) - i += 1 - test_range.append(r) - - return test_range -def createReturnsFromList(returns, start): - current = start.replace(tzinfo=pytz.utc) - one_day = datetime.timedelta(days = 1) - test_range = [] - i = 0 - while len(test_range) < len(returns): - if(risk.trading_calendar.is_trading_day(current)): - r = daily_return(current, returns[i]) - i += 1 - test_range.append(r) - current = current + one_day - return test_range - - -def createAlgo(filename): - algo = Algorithm() - algo.code = getCodeFromFile(filename) - algo.title = filename - algo._id = pymongo.objectid.ObjectId() - hostedAlgo = HostedAlgorithm(algo) - return hostedAlgo - -def getCodeFromFile(filename): - rVal = None - with open('./test/algo_samples/' + filename, 'r') as f: - rVal = f.read() - return rVal - - + def create_trade(sid, price, amount, datetime): row = {} row['source_id'] = "test_factory" - row['type'] = "TRADE" + row['type'] = zp.DATASOURCE_TYPE.TRADE row['sid'] = sid row['dt'] = datetime row['price'] = price diff --git a/zipline/test/test_finance.py b/zipline/test/test_finance.py index 5bb3fdfa..c95b2b52 100644 --- a/zipline/test/test_finance.py +++ b/zipline/test/test_finance.py @@ -13,7 +13,7 @@ import zipline.protocol as zp from zipline.test.client import TestTradingClient from zipline.test.dummy import ThreadPoolExecutorMixin from zipline.sources import SpecificEquityTrades -from zipline.finance.trading import TradeSimulator +from zipline.finance.trading import TransactionSimulator, OrderDataSource class FinanceTestCase(ThreadPoolExecutorMixin, TestCase): @@ -36,7 +36,7 @@ class FinanceTestCase(ThreadPoolExecutorMixin, TestCase): #do a transform trans_msg = zp.TRANSFORM_FRAME('helloworld', 2345.6) #simulate passthrough transform -- passthrough shouldn't even unpack the msg, just resend. - passthrough_msg = zp.TRANSFORM_FRAME('PASSTHROUGH', feed_msg) + passthrough_msg = zp.TRANSFORM_FRAME(zp.TRANSFORM_TYPE.PASSTHROUGH, feed_msg) #merge unframes transform and passthrough trans_recovered = zp.TRANSFORM_UNFRAME(trans_msg) pt_recovered = zp.TRANSFORM_UNFRAME(passthrough_msg) @@ -54,11 +54,48 @@ class FinanceTestCase(ThreadPoolExecutorMixin, TestCase): self.assertEqual(zp.namedict(trade), event) def test_order_protocol(self): + #client places an order order_msg = zp.ORDER_FRAME(133, 100) + + #order datasource receives sid, amount = zp.ORDER_UNFRAME(order_msg) self.assertEqual(sid, 133) self.assertEqual(amount, 100) - + + #order datasource datasource frames the order + order_dt = datetime.datetime.utcnow().replace(tzinfo=pytz.utc) + order_event = zp.namedict({"sid" : sid, + "amount" : amount, + "dt" : order_dt, + "source_id" : zp.FINANCE_COMPONENT.ORDER_SOURCE, + "type" : zp.DATASOURCE_TYPE.ORDER + }) + + order_ds_msg = zp.DATASOURCE_FRAME(order_event) + + #transaction transform unframes + recovered_order = zp.DATASOURCE_UNFRAME(order_ds_msg) + + self.assertEqual(order_dt, recovered_order.dt) + + #create a transaction from the order + txn = zp.namedict({ + 'sid' : recovered_order.sid, + 'amount' : recovered_order.amount, + 'dt' : recovered_order.dt, + 'price' : 10.0, + 'commission' : 0.50 + }) + + #frame that transaction + txn_msg = zp.TRANSFORM_FRAME(zp.TRANSFORM_TYPE.TRANSACTION, txn) + + #unframe + recovered_tx = zp.TRANSFORM_UNFRAME(txn_msg).TRANSACTION + self.assertEqual(recovered_tx.sid, 133) + self.assertEqual(recovered_tx.amount, 100) + + def test_trading_calendar(self): known_trading_day = datetime.datetime.strptime("02/24/2012","%m/%d/%Y") known_holiday = datetime.datetime.strptime("02/20/2012", "%m/%d/%Y") #president's day @@ -95,10 +132,12 @@ class FinanceTestCase(ThreadPoolExecutorMixin, TestCase): [100,100,100,100,100,100,100,100,100,100,100,100,100,100,100,100], datetime.datetime.strptime("02/1/2012","%m/%d/%Y"), datetime.timedelta(days=1))) - client = TestTradingClient(10) - order_sim = TradeSimulator(expected_orders=10) - - sim.register_components([client, order_sim, set1]) + #client sill send 10 orders for 100 shares of 133 + client = TestTradingClient(133, 100, 10) + order_source = OrderDataSource(datetime.datetime.strptime("02/1/2012","%m/%d/%Y").replace(tzinfo=pytz.utc)) + transaction_sim = TransactionSimulator() + + sim.register_components([client, order_source, transaction_sim, set1]) sim.register_controller( con ) # Simulation