protocol tests passing, orders test almost passing, but looking like we need to extend the merge component to add transaction and order specific logic.

This commit is contained in:
fawce
2012-02-29 23:38:46 -05:00
parent 00c2ccfe72
commit 536a1e7fdc
6 changed files with 211 additions and 158 deletions
+59 -45
View File
@@ -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)
+3 -2
View File
@@ -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):
+93 -38
View File
@@ -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'
})
+7 -9
View File
@@ -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
+3 -57
View File
@@ -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
+46 -7
View File
@@ -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