mirror of
https://github.com/wassname/catalyst.git
synced 2026-08-01 12:20:21 +08:00
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:
+59
-45
@@ -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)
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -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
@@ -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'
|
||||
})
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user