From 00c2ccfe72c345809e980ec990ea1b0465f64092 Mon Sep 17 00:00:00 2001 From: fawce Date: Wed, 29 Feb 2012 17:18:17 -0500 Subject: [PATCH 1/3] using an enum for finance protocol constants --- zipline/finance/trading.py | 10 +++++----- zipline/protocol.py | 6 +++--- 2 files changed, 8 insertions(+), 8 deletions(-) diff --git a/zipline/finance/trading.py b/zipline/finance/trading.py index ee13d432..e3a97e53 100644 --- a/zipline/finance/trading.py +++ b/zipline/finance/trading.py @@ -71,12 +71,12 @@ class OrderDataSource(qmsg.DataSource): 'volume' : integer for volume } """ - zm.DataSource.__init__(self, "ORDER_SIM") + zm.DataSource.__init__(self, str(zp.FINANCE_PROTOCOL.ORDER)) self.simulation_dt = simulation_dt self.last_iteration_duration = datetime.timedelta(seconds=0) def get_type(self): - return 'ORDER_SIM' + return str(zp.FINANCE_PROTOCOL.ORDER) def open(self): qmsg.DataSource.open(self) @@ -88,7 +88,7 @@ class OrderDataSource(qmsg.DataSource): 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('ORDER_SIM', self.simulation_dt), self.zmq.NOBLOCK) + 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 @@ -139,10 +139,10 @@ class TransactionSimulator(qmsg.BaseTransform): Pulls one message from the event feed, then loops on orders until client sends DONE message. """ - if(event.type == "ORDER_SIM"): + if(event.type == zp.FINANCE_PROTOCOL.ORDER): self.add_open_order(event.sid, event.amount) self.state['value'] = self.average - elif(event.type == "EQUITY_TRADE"): + elif(event.type == zp.FINANCE_PROTOCOL.TRADE): txn = apply_trade_to_open_orders(event) self.state['value'] = txn diff --git a/zipline/protocol.py b/zipline/protocol.py index dc3fb408..22ca9712 100644 --- a/zipline/protocol.py +++ b/zipline/protocol.py @@ -429,7 +429,7 @@ def UNPACK_DATE(payload): FINANCE_PROTOCOL = Enum( - 'ORDER' , # 0 - req - 'TRANSACTION' , # 1 - req - 'STATUS' , # 2 - req + 'ORDER' , # 0 + 'TRANSACTION' , # 1 + 'TRADE' , # 2 ) From 536a1e7fdc9f655883921da3e3ff63f902e88176 Mon Sep 17 00:00:00 2001 From: fawce Date: Wed, 29 Feb 2012 23:38:46 -0500 Subject: [PATCH 2/3] 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 From 7846a1e1f49f3754f137617c99e6800438ef70f3 Mon Sep 17 00:00:00 2001 From: fawce Date: Thu, 1 Mar 2012 10:46:50 -0500 Subject: [PATCH 3/3] passing basic order and transaction simulation test --- zipline/finance/trading.py | 112 +++++++++++++++++++---------------- zipline/protocol.py | 34 ++++++++--- zipline/test/client.py | 5 +- zipline/test/test_finance.py | 13 ++-- 4 files changed, 99 insertions(+), 65 deletions(-) diff --git a/zipline/finance/trading.py b/zipline/finance/trading.py index a6b81d2e..9420eb96 100644 --- a/zipline/finance/trading.py +++ b/zipline/finance/trading.py @@ -1,5 +1,7 @@ import json import datetime +import pytz +import math from zmq.core.poll import select @@ -35,25 +37,26 @@ class TradeSimulationClient(qmsg.Component): self.signal_done() return - event = zp.MERGE_UNFRAME(message) + event = zp.MERGE_UNFRAME(msg) self._handle_event(event) def connect_order(self): return self.connect_push_socket(self.addresses['order_address']) def _handle_event(self, event): - self.event_queue.append(event) - if event.SIM_DT <= event.dt: - #event occurred in the present, send the queue to be processed - self.handle_events(self.event_queue) - self.order_socket.send(str(zp.CONTROL_PROTOCOL.DONE)) + self.handle_event(event) + #signal done to order source. + self.order_socket.send(str(zp.ORDER_PROTOCOL.BREAK)) - def handle_events(self, event_queue): + def handle_event(self, event): raise NotImplementedError def order(self, sid, amount): self.order_socket.send(zp.ORDER_FRAME(sid, amount)) + def signal_order_done(self): + self.order_socket.send(str(zp.ORDER_PROTOCOL.DONE)) + class OrderDataSource(qmsg.DataSource): """DataSource that relays orders from the client""" @@ -107,7 +110,11 @@ class OrderDataSource(qmsg.DataSource): continue order_msg = rlist[0].recv() - if order_msg == str(zp.CONTROL_PROTOCOL.DONE): + if order_msg == str(zp.ORDER_PROTOCOL.DONE): + self.signal_done() + return + + if order_msg == str(zp.ORDER_PROTOCOL.BREAK): qutil.LOGGER.info("order loop finished") break @@ -156,74 +163,77 @@ class TransactionSimulator(qmsg.BaseTransform): """ #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.create_transaction(0, 0, 0.0, event.dt, 1) + self.add_open_order(event) + self.state['value'] = None elif(event.type == zp.DATASOURCE_TYPE.TRADE): - txn = apply_trade_to_open_orders(event) + txn = self.apply_trade_to_open_orders(event) self.state['value'] = txn + else: + self.state['value'] = None + qutil.LOGGER.info("unexpected event type in transform: {etype}".format(etype=event.type)) #TODO: what to do if we get another kind of datasource event.type? return self.state - def add_open_order(self, sid, amount): + def add_open_order(self, event): """Orders are captured in a buffer by sid. No calculations are done here. Amount is explicitly converted to an int. Orders of amount zero are ignored. """ - amount = int(amount) - if amount == 0: - qutil.LOGGER.debug("requested to trade zero shares of {sid}".format(sid=sid)) + event.amount = int(event.amount) + if event.amount == 0: + qutil.LOGGER.debug("requested to trade zero shares of {sid}".format(sid=event.sid)) return self.order_count += 1 - order = zp.namedict({'sid' : sid, - 'amount' : amount, - 'dt' : self.algo_time, - 'filled': 0, - 'direction': math.fabs(amount) / amount - }) - if(not self.open_orders.has_key(sid)): - self.open_orders[sid] = [] - self.open_orders[sid].append(order) + if(not self.open_orders.has_key(event.sid)): + self.open_orders[event.sid] = [] + self.open_orders[event.sid].append(event) def apply_trade_to_open_orders(self, event): if(event.volume == 0): #there are zero volume events bc some stocks trade less frequently than once per minute. - return + return self.create_dummy_txn(event.dt) + if self.open_orders.has_key(event.sid): orders = self.open_orders[event.sid] - remaining_orders = [] - total_order = 0 - dt = event.dt - - for order in orders: - #we're using minute bars, so allow orders within 30 seconds of the trade - if((order.dt - event.dt) < self.trade_windwo): - total_order += order.amount - if(order.dt > dt): - dt = order.dt - #if the order still has time to live (TTL) keep track - elif((self.algo_time - order.dt) < self.orderTTL): - remaining_orders.append(order) - - self.open_orders[event.sid] = remaining_orders - - if(total_order != 0): - direction = total_order / math.fabs(total_order) - volume_share = (direction * total_order) / event.volume - if volume_share > .25: - volume_share = .25 - amount = volume_share * event.volume * direction - impact = (volShare)**2 * .1 * direction * event.price - return self.create_transaction(event.sid, amount, event.price + impact, dt.replace(tzinfo = pytz.utc), direction) - + else: + return None + + remaining_orders = [] + total_order = 0 + dt = event.dt + + for order in orders: + #we're using minute bars, so allow orders within 30 seconds of the trade + if((order.dt - event.dt) < self.trade_windwo): + total_order += order.amount + if(order.dt > dt): + dt = order.dt + #if the order still has time to live (TTL) keep track + elif((self.algo_time - order.dt) < self.orderTTL): + remaining_orders.append(order) + + self.open_orders[event.sid] = remaining_orders + + if(total_order != 0): + direction = total_order / math.fabs(total_order) + else: + direction = 1 + + volume_share = (direction * total_order) / event.volume + if volume_share > .25: + volume_share = .25 + amount = volume_share * event.volume * direction + impact = (volume_share)**2 * .1 * direction * event.price + return self.create_transaction(event.sid, amount, event.price + impact, dt.replace(tzinfo = pytz.utc), direction) def create_transaction(self, sid, amount, price, dt, direction): txn = {'sid' : sid, - 'amount' : amount, + 'amount' : int(amount), 'dt' : dt, 'price' : price, 'commission' : self.commission * amount * direction, diff --git a/zipline/protocol.py b/zipline/protocol.py index fe074840..e2b9bae2 100644 --- a/zipline/protocol.py +++ b/zipline/protocol.py @@ -106,7 +106,7 @@ class namedict(object): return "namedict: " + str(self.__dict__) def __eq__(self, other): - return self.__dict__ == other.__dict__ + return other != None and self.__dict__ == other.__dict__ def has_attr(self, name): return self.__dict__.has_key(name) @@ -263,13 +263,9 @@ def FEED_UNFRAME(msg): INVALID_TRANSFORM_FRAME = FrameExceptionFactory('TRANSFORM') def TRANSFORM_FRAME(name, value): - """ - :event: a nameddict with at least:: - - source_id - - type - """ assert isinstance(name, basestring) - assert value != None + if value == None: + return msgpack.dumps(tuple([name, TRANSFORM_TYPE.EMPTY])) if(name == TRANSFORM_TYPE.TRANSACTION): value = TRANSACTION_FRAME(value) return msgpack.dumps(tuple([name, value])) @@ -279,13 +275,17 @@ def TRANSFORM_UNFRAME(msg): :rtype: namedict with : """ try: + name, value = msgpack.loads(msg) + if(value == TRANSFORM_TYPE.EMPTY): + return namedict({name : None}) #TODO: anything we can do to assert more about the content of the dict? assert isinstance(name, basestring) if(name == TRANSFORM_TYPE.PASSTHROUGH): value = FEED_UNFRAME(value) elif(name == TRANSFORM_TYPE.TRANSACTION): value = TRANSACTION_UNFRAME(value) + return namedict({name : value}) except TypeError: raise INVALID_TRANSFORM_FRAME(msg) @@ -305,6 +305,11 @@ def MERGE_FRAME(event): """ assert isinstance(event, namedict) PACK_DATE(event) + if(event.has_attr(TRANSFORM_TYPE.TRANSACTION)): + if(event.TRANSACTION == None): + event.TRANSACTION = TRANSFORM_TYPE.EMPTY + else: + event.TRANSACTION = TRANSACTION_FRAME(event.TRANSACTION) payload = event.__dict__ return msgpack.dumps(payload) @@ -314,6 +319,11 @@ 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(TRANSFORM_TYPE.TRANSACTION)): + if(payload.TRANSACTION == TRANSFORM_TYPE.EMPTY): + payload.TRANSACTION = None + else: + payload.TRANSACTION = TRANSACTION_UNFRAME(payload.TRANSACTION) UNPACK_DATE(payload) return payload except TypeError: @@ -476,9 +486,17 @@ DATASOURCE_TYPE = Enum( 'TRADE' , ) +ORDER_PROTOCOL = Enum( + 'DONE', + 'BREAK' +) + + +#Transform type needs to be a namedict to facilitate merging. TRANSFORM_TYPE = namedict({ 'TRANSACTION':'TRANSACTION', #needed? - 'PASSTHROUGH':'PASSTHROUGH' + 'PASSTHROUGH':'PASSTHROUGH', + 'EMPTY':'' }) diff --git a/zipline/test/client.py b/zipline/test/client.py index 02a8e4c1..00259118 100644 --- a/zipline/test/client.py +++ b/zipline/test/client.py @@ -51,9 +51,12 @@ class TestTradingClient(TradeSimulationClient): self.amount = amount self.incr = 0 - def handle_events(self, event_queue): + def handle_event(self, event): #place an order for 100 shares of sid:133 if(self.incr < self.count): self.order(self.sid, self.amount) self.incr += 1 + else: + self.signal_order_done() + self.signal_done() diff --git a/zipline/test/test_finance.py b/zipline/test/test_finance.py index c95b2b52..cd5b0a32 100644 --- a/zipline/test/test_finance.py +++ b/zipline/test/test_finance.py @@ -127,11 +127,14 @@ class FinanceTestCase(ThreadPoolExecutorMixin, TestCase): # Simulation Components # --------------------- - set1 = SpecificEquityTrades("flat-133",factory.create_trade_history(133, - [10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0], - [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))) + set1 = SpecificEquityTrades("flat-133", + factory.create_trade_history( + 133, + [10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0,10.0], + [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 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))