Merge branch 'finance' into master_finance_merge

Conflicts:
	zipline/messaging.py
	zipline/test/client.py
This commit is contained in:
fawce
2012-03-01 13:54:47 -05:00
9 changed files with 330 additions and 350 deletions
+1 -1
View File
@@ -3,4 +3,4 @@ pyzmq==2.1.11
gevent-zeromq==0.2.2
msgpack-python==0.1.12
humanhash==0.0.1
ujson=1.18
ujson==1.18
+2 -2
View File
@@ -12,6 +12,6 @@ detailed-errors=1
# Drop into debugger on failure
pdb=0
pdb-failures=0
#pdb=0
#pdb-failures=0
+111 -87
View File
@@ -1,5 +1,7 @@
import json
import datetime
import pytz
import math
from zmq.core.poll import select
@@ -17,7 +19,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,38 +27,36 @@ 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(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"""
@@ -71,12 +71,13 @@ class OrderDataSource(qmsg.DataSource):
'volume' : integer for volume
}
"""
zm.DataSource.__init__(self, "ORDER_SIM")
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 'ORDER_SIM'
return zp.FINANCE_COMPONENT.ORDER_SOURCE
def open(self):
qmsg.DataSource.open(self)
@@ -85,16 +86,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('ORDER_SIM', 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],
[],
@@ -106,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
@@ -115,21 +123,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,83 +161,85 @@ 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"):
self.add_open_order(event.sid, event.amount)
self.state['value'] = self.average
elif(event.type == "EQUITY_TRADE"):
txn = apply_trade_to_open_orders(event)
#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)
self.state['value'] = None
elif(event.type == zp.DATASOURCE_TYPE.TRADE):
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("{title}:{id} requested to trade zero shares of {sid}".format(sid=sid,
title=self.hostedAlgo.algo.title,
id=self.hostedAlgo.algo.id))
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.
continue
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.tradeWindow):
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):
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' : int(amount),
'dt' : dt,
'price' : price,
'commission' : self.commission * amount * direction,
'source_id' : zp.FINANCE_COMPONENT.TRANSACTION_SIM
}
return zp.namedict(txn)
+2 -2
View File
@@ -476,7 +476,6 @@ class PassthroughTransform(BaseTransform):
def __init__(self):
BaseTransform.__init__(self, "PASSTHROUGH")
self.init()
def init(self):
@@ -486,8 +485,9 @@ class PassthroughTransform(BaseTransform):
def get_type(self):
return COMPONENT_TYPE.CONDUIT
#TODO, could save some cycles by skipping the _UNFRAME call and just setting value to original msg string.
def transform(self, event):
return { 'value': event }
return {'name':zp.TRANSFORM_TYPE.PASSTHROUGH, 'value': zp.DATASOURCE_FRAME(event) }
class DataSource(Component):
+118 -45
View File
@@ -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)
@@ -196,9 +196,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))
@@ -217,9 +219,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)
@@ -265,17 +269,11 @@ 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(name == 'SIM_DT'):
value = PACK_ALGO_DT(value)
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]))
def TRANSFORM_UNFRAME(msg):
@@ -283,29 +281,23 @@ def TRANSFORM_UNFRAME(msg):
:rtype: namedict with <transform_name>:<transform_value>
"""
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 == "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
# ==================
@@ -318,10 +310,12 @@ 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)
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)
@@ -331,10 +325,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('SIM_DT')):
payload.SIM_DT = UNPACK_ALGO_DT(payload.SIM_DT)
assert isinstance(payload.epoch, numbers.Integral)
assert isinstance(payload.micros, numbers.Integral)
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:
@@ -350,7 +345,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):
@@ -364,7 +359,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)
@@ -378,8 +373,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
@@ -389,7 +382,7 @@ def TRADE_UNFRAME(msg):
raise INVALID_TRADE_FRAME(msg)
# =========
# Orders
# Orders - from client to order source
# =========
def ORDER_FRAME(sid, amount):
@@ -410,6 +403,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
# =================
@@ -433,8 +486,28 @@ def UNPACK_DATE(payload):
payload['dt'] = dt
return payload
FINANCE_PROTOCOL = Enum(
'ORDER' , # 0 - req
'TRANSACTION' , # 1 - req
'STATUS' , # 2 - req
)
DATASOURCE_TYPE = Enum(
'ORDER' ,
'TRADE' ,
)
ORDER_PROTOCOL = Enum(
'DONE',
'BREAK'
)
#Transform type needs to be a namedict to facilitate merging.
TRANSFORM_TYPE = namedict({
'TRANSACTION':'TRANSACTION', #needed?
'PASSTHROUGH':'PASSTHROUGH',
'EMPTY':''
})
FINANCE_COMPONENT = namedict({
'TRADING_CLIENT':'TRADING_CLIENT',
'PORTFOLIO_CLIENT':'PORTFOLIO_CLIENT',
'ORDER_SOURCE':'ORDER_SOURCE',
'TRANSACTION_SIM':'TRANSACTION_SIM'
})
+22 -9
View File
@@ -4,22 +4,17 @@ import ujson as json
import zipline.util as qutil
import zipline.messaging as qmsg
from zipline.protocol import CONTROL_PROTOCOL, COMPONENT_TYPE
#qlogger = logging.getLogger('qexec')
from zipline.finance.trading import TradeSimulationClient
class TestClient(qmsg.Component):
def __init__(self, expected_msg_count=0):
def __init__(self):
qmsg.Component.__init__(self)
# Generally shouldn't reference unit tests here.
self.utest = utest
self.expected_msg_count = expected_msg_count
self.init()
def init(self):
self.received_count = 0
self.prev_dt = None
self.received_count = 0
self.prev_dt = None
@property
def get_id(self):
@@ -68,3 +63,21 @@ class TestClient(qmsg.Component):
if self.received_count % 100 == 0:
qutil.LOGGER.info("received {n} messages".format(n=self.received_count))
class TestTradingClient(TradeSimulationClient):
def __init__(self, sid, amount, order_count):
TradeSimulationClient.__init__(self)
self.count = order_count
self.sid = sid
self.amount = amount
self.incr = 0
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()
-128
View File
@@ -1,128 +0,0 @@
"""
Dummy simulator for test/development on Zipline.
"""
import threading
import mock
from collections import defaultdict
from zipline.monitor import Controller
from zipline.messaging import SimulatorBase
import zipline.util as qutil
class DummyAllocator(object):
def __init__(self, ns):
self.idx = 0
self.sockets = [
'tcp://127.0.0.1:%s' % (10000 + n)
for n in xrange(ns)
]
def lease(self, n):
sockets = self.sockets[self.idx:self.idx+n]
self.idx += n
return sockets
def reaquire(self, *conn):
pass
class ThreadSimulator(SimulatorBase):
def __init__(self, addresses):
SimulatorBase.__init__(self, addresses)
def launch_controller(self):
thread = threading.Thread(target=self.controller.run)
thread.start()
self.cuc = thread
return thread
def launch_component(self, component):
thread = threading.Thread(target=component.run)
thread.start()
return thread
class ExecutorMixinBase(object):
"""Abstract base to allow mixin for tests that need a dummy simulator."""
leased_sockets = defaultdict(list)
def setUp(self):
self.setup_logging()
# TODO: how to make Nose use this cross-process????
self.setup_allocator()
def tearDown(self):
pass
#self.unallocate_sockets()
# Assert the sockets were properly cleaned up
#self.assertEmpty(self.leased_sockets[self.id()].values())
# Assert they were returned to the heap
#self.allocator.socketheap.assert
def get_simulator(self):
"""
Return a new simulator instance to be tested.
"""
raise NotImplementedError
def get_controller(self):
"""
Return a new controler for simulator instance to be tested.
"""
raise NotImplementedError
def setup_allocator(self):
"""
Setup the socket allocator for this test case.
"""
raise NotImplementedError
def allocate_sockets(self, n):
"""
Allocate sockets local to this test case, track them so
we can gc after test run.
"""
assert isinstance(n, int)
assert n > 0
leased = self.allocator.lease(n)
self.leased_sockets[self.id()].extend(leased)
return leased
def unallocate_sockets(self):
self.allocator.reaquire(*self.leased_sockets[self.id()])
class ThreadPoolExecutorMixin(ExecutorMixinBase):
"""Dummy server using threads."""
allocator = DummyAllocator(100)
def setup_logging(self):
qutil.configure_logging()
# lazy import by design
self.logger = mock.Mock()
def setup_allocator(self):
pass
def get_simulator(self, addresses):
return ThreadSimulator(addresses)
def get_controller(self):
# Allocate two more sockets
controller_sockets = self.allocate_sockets(2)
return Controller(
controller_sockets[0],
controller_sockets[1],
logging = self.logger,
)
+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
+71 -19
View File
@@ -11,12 +11,16 @@ import zipline.finance.risk as risk
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
from zipline.simulator import AddressAllocator, Simulator
from zipline.monitor import Controller
class FinanceTestCase(ThreadPoolExecutorMixin, TestCase):
class FinanceTestCase(TestCase):
def setUp(self):
qutil.configure_logging()
def test_trade_feed_protocol(self):
trades = factory.create_trade_history(133,
@@ -36,7 +40,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 +58,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
@@ -73,8 +114,9 @@ class FinanceTestCase(ThreadPoolExecutorMixin, TestCase):
# --------------
# Allocate sockets for the simulator components
sockets = self.allocate_sockets(6)
allocator = AddressAllocator(8)
sockets = allocator.lease(8)
addresses = {
'sync_address' : sockets[0],
'data_address' : sockets[1],
@@ -83,27 +125,37 @@ class FinanceTestCase(ThreadPoolExecutorMixin, TestCase):
'result_address' : sockets[4],
'order_address' : sockets[5]
}
con = Controller(
sockets[6],
sockets[7],
logging = qutil.LOGGER
)
sim = self.get_simulator(addresses)
con = self.get_controller()
sim = Simulator(addresses)
# 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)))
client = TestTradingClient(10)
order_sim = TradeSimulator(expected_orders=10)
sim.register_components([client, order_sim, set1])
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))
transaction_sim = TransactionSimulator()
sim.register_components([client, order_source, transaction_sim, set1])
sim.register_controller( con )
# Simulation
# ----------
sim.simulate()
sim.simulate().join()
# Stop Running
# ------------