mirror of
https://github.com/wassname/catalyst.git
synced 2026-07-24 13:00:57 +08:00
Merge branch 'refactor'
Conflicts: zipline/finance/trading.py zipline/lines.py
This commit is contained in:
@@ -0,0 +1 @@
|
||||
# TODO: move qexec console here
|
||||
+37
@@ -0,0 +1,37 @@
|
||||
[loggers]
|
||||
keys=root
|
||||
|
||||
[handlers]
|
||||
keys=consoleHandler,filesystemHandler
|
||||
|
||||
[formatters]
|
||||
keys=ziplineformat
|
||||
|
||||
# -------
|
||||
|
||||
[logger_root]
|
||||
level=DEBUG
|
||||
handlers=consoleHandler,filesystemHandler
|
||||
qualname=ZiplineLogger
|
||||
|
||||
# -------
|
||||
|
||||
[handler_filesystemHandler]
|
||||
class=handlers.RotatingFileHandler
|
||||
level=DEBUG
|
||||
formatter=ziplineformat
|
||||
args=("/var/log/zipline/zipline.log",10*1024*1024,5)
|
||||
propagate=1
|
||||
|
||||
[handler_consoleHandler]
|
||||
class=StreamHandler
|
||||
level=ERROR
|
||||
formatter=ziplineformat
|
||||
args=(sys.stdout,)
|
||||
propagate=1
|
||||
|
||||
# -------
|
||||
|
||||
[formatter_ziplineformat]
|
||||
format=%(asctime)s %(levelname)s %(filename)s %(funcName)s - %(message)s
|
||||
datefmt=%Y-%m-%d %H:%M:%S %Z
|
||||
@@ -1,15 +1,16 @@
|
||||
import logging
|
||||
from gevent_zeromq import zmq
|
||||
|
||||
import zipline.util as qutil
|
||||
import zipline.messaging as qmsg
|
||||
import zipline.protocol as zp
|
||||
from zipline.core.component import Component
|
||||
from zipline.protocol import CONTROL_PROTOCOL, COMPONENT_TYPE
|
||||
from zipline.finance.trading import TradeSimulationClient
|
||||
|
||||
class TestClient(qmsg.Component):
|
||||
LOGGER = logging.getLogger('ZiplineLogger')
|
||||
|
||||
class TestClient(Component):
|
||||
|
||||
def __init__(self):
|
||||
qmsg.Component.__init__(self)
|
||||
Component.__init__(self)
|
||||
self.init()
|
||||
|
||||
def init(self):
|
||||
@@ -55,7 +56,7 @@ class TestClient(qmsg.Component):
|
||||
#logger.info('msg:' + str(msg))
|
||||
|
||||
if msg == str(CONTROL_PROTOCOL.DONE):
|
||||
qutil.LOGGER.info("Client is DONE!")
|
||||
LOGGER.info("Client is DONE!")
|
||||
self.signal_done()
|
||||
return
|
||||
|
||||
@@ -79,7 +80,7 @@ class TestClient(qmsg.Component):
|
||||
self.prev_dt = event.dt
|
||||
|
||||
if self.received_count % 100 == 0:
|
||||
qutil.LOGGER.info("received {n} messages".format(n=self.received_count))
|
||||
|
||||
LOGGER.info("received {n} messages".format(n=self.received_count))
|
||||
|
||||
def unframe(self, msg):
|
||||
return zp.MERGE_UNFRAME(msg)
|
||||
@@ -1,5 +1,6 @@
|
||||
"""Tests for the zipline.finance package"""
|
||||
import mock
|
||||
"""
|
||||
Tests for the zipline.finance package
|
||||
"""
|
||||
import pytz
|
||||
|
||||
from unittest2 import TestCase
|
||||
@@ -8,22 +9,16 @@ from collections import defaultdict
|
||||
|
||||
from nose.tools import timed
|
||||
|
||||
import zipline.test.factory as factory
|
||||
import zipline.util as qutil
|
||||
import zipline.finance.risk as risk
|
||||
import zipline.utils.factory as factory
|
||||
import zipline.protocol as zp
|
||||
import zipline.finance.performance as perf
|
||||
|
||||
from zipline.test.algorithms import TestAlgorithm
|
||||
from zipline.sources import SpecificEquityTrades
|
||||
from zipline.finance.trading import TransactionSimulator, \
|
||||
TradeSimulationClient, TradingEnvironment
|
||||
from zipline.simulator import AddressAllocator, Simulator
|
||||
from zipline.monitor import Controller
|
||||
from zipline.test_algorithms import TestAlgorithm
|
||||
from zipline.finance.trading import TradingEnvironment
|
||||
from zipline.core.devsimulator import AddressAllocator
|
||||
from zipline.lines import SimulatedTrading
|
||||
from zipline.finance.performance import PerformanceTracker
|
||||
from zipline.protocol_utils import namedict
|
||||
from zipline.finance.trading import SIMULATION_STYLE
|
||||
from zipline.utils.protocol_utils import ndict
|
||||
from zipline.finance.trading import TransactionSimulator, SIMULATION_STYLE
|
||||
|
||||
DEFAULT_TIMEOUT = 15 # seconds
|
||||
EXTENDED_TIMEOUT = 90
|
||||
@@ -35,7 +30,7 @@ class FinanceTestCase(TestCase):
|
||||
leased_sockets = defaultdict(list)
|
||||
|
||||
def setUp(self):
|
||||
qutil.configure_logging()
|
||||
#qutil.configure_logging()
|
||||
self.zipline_test_config = {
|
||||
'allocator':allocator,
|
||||
'sid':133
|
||||
@@ -148,7 +143,7 @@ class FinanceTestCase(TestCase):
|
||||
|
||||
# TODO: for some reason the orders aren't filled without an extra
|
||||
# trade.
|
||||
trade_count = 5001
|
||||
trade_count = 5
|
||||
self.zipline_test_config['order_count'] = trade_count - 1
|
||||
self.zipline_test_config['trade_count'] = trade_count
|
||||
self.zipline_test_config['order_amount'] = 1
|
||||
@@ -156,7 +151,7 @@ class FinanceTestCase(TestCase):
|
||||
# tell the simulator to fill the orders in individual transactions
|
||||
# matching the order volume exactly.
|
||||
self.zipline_test_config['simulation_style'] = \
|
||||
SIMULATION_STYLE.FIXED_SLIPPAGE
|
||||
SIMULATION_STYLE.FIXED_SLIPPAGE
|
||||
self.zipline_test_config['environment'] = factory.create_trading_environment()
|
||||
|
||||
sid_list = [self.zipline_test_config['sid']]
|
||||
@@ -416,7 +411,7 @@ class FinanceTestCase(TestCase):
|
||||
alternate = params.get('alternate')
|
||||
# if present, expect transaction amounts to match orders exactly.
|
||||
complete_fill = params.get('complete_fill')
|
||||
|
||||
|
||||
trading_environment = factory.create_trading_environment()
|
||||
trade_sim = TransactionSimulator()
|
||||
price = [10.1] * trade_count
|
||||
@@ -424,22 +419,22 @@ class FinanceTestCase(TestCase):
|
||||
start_date = trading_environment.first_open
|
||||
sid = 1
|
||||
|
||||
generated_trades = factory.create_trade_history(
|
||||
sid,
|
||||
price,
|
||||
volume,
|
||||
trade_interval,
|
||||
trading_environment
|
||||
generated_trades = factory.create_trade_history(
|
||||
sid,
|
||||
price,
|
||||
volume,
|
||||
trade_interval,
|
||||
trading_environment
|
||||
)
|
||||
|
||||
|
||||
if alternate:
|
||||
alternator = -1
|
||||
else:
|
||||
alternator = 1
|
||||
|
||||
|
||||
order_date = start_date
|
||||
for i in xrange(order_count):
|
||||
order = namedict(
|
||||
order = ndict(
|
||||
{
|
||||
'sid' : sid,
|
||||
'amount' : order_amount * alternator**i,
|
||||
@@ -448,7 +443,7 @@ class FinanceTestCase(TestCase):
|
||||
})
|
||||
|
||||
trade_sim.add_open_order(order)
|
||||
|
||||
|
||||
order_date = order_date + order_interval
|
||||
# move after market orders to just after market next
|
||||
# market open.
|
||||
@@ -456,40 +451,40 @@ class FinanceTestCase(TestCase):
|
||||
if order_date.minute >= 00:
|
||||
order_date = order_date + timedelta(days=1)
|
||||
order_date = order_date.replace(hour=14, minute=30)
|
||||
|
||||
|
||||
# there should now be one open order list stored under the sid
|
||||
oo = trade_sim.open_orders
|
||||
self.assertEqual(len(oo), 1)
|
||||
self.assertTrue(oo.has_key(sid))
|
||||
order_list = oo[sid]
|
||||
self.assertEqual(order_count, len(order_list))
|
||||
|
||||
|
||||
for i in xrange(order_count):
|
||||
order = order_list[i]
|
||||
self.assertEqual(order.sid, sid)
|
||||
self.assertEqual(order.amount, order_amount * alternator**i)
|
||||
|
||||
|
||||
|
||||
|
||||
tracker = PerformanceTracker(trading_environment)
|
||||
|
||||
|
||||
# this approximates the loop inside TradingSimulationClient
|
||||
transactions = []
|
||||
for trade in generated_trades:
|
||||
if trade_delay:
|
||||
trade.dt = trade.dt + trade_delay
|
||||
|
||||
|
||||
txn = trade_sim.apply_trade_to_open_orders(trade)
|
||||
if txn:
|
||||
transactions.append(txn)
|
||||
trade.TRANSACTION = txn
|
||||
transactions.append(txn)
|
||||
trade.TRANSACTION = txn
|
||||
else:
|
||||
trade.TRANSACTION = None
|
||||
|
||||
tracker.process_event(trade)
|
||||
|
||||
|
||||
tracker.process_event(trade)
|
||||
|
||||
if complete_fill:
|
||||
self.assertEqual(len(transactions), len(order_list))
|
||||
|
||||
self.assertEqual(len(transactions), len(order_list))
|
||||
|
||||
total_volume = 0
|
||||
for i in xrange(len(transactions)):
|
||||
txn = transactions[i]
|
||||
@@ -497,18 +492,18 @@ class FinanceTestCase(TestCase):
|
||||
if complete_fill:
|
||||
order = order_list[i]
|
||||
self.assertEqual(order.amount, txn.amount)
|
||||
|
||||
self.assertEqual(total_volume, expected_txn_volume)
|
||||
|
||||
self.assertEqual(total_volume, expected_txn_volume)
|
||||
self.assertEqual(len(transactions), expected_txn_count)
|
||||
|
||||
|
||||
cumulative_pos = tracker.cumulative_performance.positions[sid]
|
||||
self.assertEqual(total_volume, cumulative_pos.amount)
|
||||
|
||||
|
||||
# the open orders should now be empty
|
||||
oo = trade_sim.open_orders
|
||||
self.assertTrue(oo.has_key(sid))
|
||||
order_list = oo[sid]
|
||||
self.assertEqual(0, len(order_list))
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from zipline.protocol_utils import ndict, namedict
|
||||
from zipline.utils.protocol_utils import ndict
|
||||
|
||||
def test_ndict():
|
||||
nd = ndict({})
|
||||
@@ -21,11 +21,18 @@ def test_ndict():
|
||||
assert 'x' in nd
|
||||
assert 'y' not in nd
|
||||
|
||||
# Mutability
|
||||
nd2 = ndict({'x': 1})
|
||||
assert nd2.x == 1
|
||||
nd2.x = 2
|
||||
assert nd2.x == 2
|
||||
|
||||
# Class isolation
|
||||
assert '__init__' not in nd
|
||||
assert '__iter__' not in nd
|
||||
assert not nd.__dict__.has_key('x')
|
||||
assert nd.get('__init__') is None
|
||||
assert 'x' not in set(dir(nd))
|
||||
|
||||
# Comparison
|
||||
nd2 = nd.copy()
|
||||
@@ -4,18 +4,19 @@ import random
|
||||
import datetime
|
||||
import pytz
|
||||
|
||||
import zipline.test.factory as factory
|
||||
import zipline.test.algorithms
|
||||
import zipline.util as qutil
|
||||
import zipline.utils.factory as factory
|
||||
import zipline.test_algorithms
|
||||
#import zipline.util as qutil
|
||||
import zipline.finance.performance as perf
|
||||
import zipline.finance.risk as risk
|
||||
import zipline.protocol as zp
|
||||
from zipline.finance.trading import TradeSimulationClient, TradingEnvironment, \
|
||||
SIMULATION_STYLE
|
||||
SIMULATION_STYLE
|
||||
|
||||
class PerformanceTestCase(unittest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
qutil.configure_logging()
|
||||
#qutil.configure_logging()
|
||||
self.benchmark_returns, self.treasury_curves = \
|
||||
factory.load_market_data()
|
||||
|
||||
@@ -546,7 +547,7 @@ shares in position"
|
||||
#create a transaction for all but
|
||||
#first trade in each sid, to simulate None transaction
|
||||
if(event.dt != self.trading_environment.period_start):
|
||||
txn = zp.namedict({
|
||||
txn = zp.ndict({
|
||||
'sid' : event.sid,
|
||||
'amount' : -25,
|
||||
'dt' : event.dt,
|
||||
@@ -565,4 +566,4 @@ shares in position"
|
||||
cumulative_pos = perf_tracker.cumulative_performance.positions[sid]
|
||||
expected_size = txn_count / 2 * -25
|
||||
self.assertEqual(cumulative_pos.amount, expected_size)
|
||||
|
||||
|
||||
@@ -9,11 +9,11 @@ from collections import defaultdict
|
||||
|
||||
from nose.tools import timed
|
||||
|
||||
import zipline.test.factory as factory
|
||||
import zipline.util as qutil
|
||||
import zipline.utils.factory as factory
|
||||
from zipline.utils import logger
|
||||
import zipline.protocol as zp
|
||||
|
||||
from zipline.sources import SpecificEquityTrades
|
||||
from zipline.finance.sources import SpecificEquityTrades
|
||||
|
||||
DEFAULT_TIMEOUT = 5 # seconds
|
||||
|
||||
@@ -22,7 +22,7 @@ class ProtocolTestCase(TestCase):
|
||||
leased_sockets = defaultdict(list)
|
||||
|
||||
def setUp(self):
|
||||
qutil.configure_logging()
|
||||
#qutil.configure_logging()
|
||||
self.trading_environment = factory.create_trading_environment()
|
||||
|
||||
@timed(DEFAULT_TIMEOUT)
|
||||
@@ -45,7 +45,7 @@ class ProtocolTestCase(TestCase):
|
||||
|
||||
for trade in trades:
|
||||
#simulate data source sending frame
|
||||
msg = zp.DATASOURCE_FRAME(zp.namedict(trade))
|
||||
msg = zp.DATASOURCE_FRAME(zp.ndict(trade))
|
||||
#feed unpacking frame
|
||||
recovered_trade = zp.DATASOURCE_UNFRAME(msg)
|
||||
#feed sending frame
|
||||
@@ -74,13 +74,13 @@ class ProtocolTestCase(TestCase):
|
||||
self.assertTrue(event.helloworld == 2345.6)
|
||||
event.delete('helloworld')
|
||||
|
||||
self.assertEqual(zp.namedict(trade), event)
|
||||
self.assertEqual(zp.ndict(trade), event)
|
||||
|
||||
@timed(DEFAULT_TIMEOUT)
|
||||
def test_order_protocol(self):
|
||||
#client places an order
|
||||
now = datetime.utcnow().replace(tzinfo=pytz.utc)
|
||||
order = zp.namedict({
|
||||
order = zp.ndict({
|
||||
'dt':now,
|
||||
'sid':133,
|
||||
'amount':100
|
||||
@@ -94,7 +94,7 @@ class ProtocolTestCase(TestCase):
|
||||
self.assertEqual(order.dt, now)
|
||||
|
||||
#order datasource datasource frames the order
|
||||
order_event = zp.namedict({
|
||||
order_event = zp.ndict({
|
||||
"sid" : order.sid,
|
||||
"amount" : order.amount,
|
||||
"dt" : order.dt,
|
||||
@@ -111,7 +111,7 @@ class ProtocolTestCase(TestCase):
|
||||
self.assertEqual(now, recovered_order.dt)
|
||||
|
||||
#create a transaction from the order
|
||||
txn = zp.namedict({
|
||||
txn = zp.ndict({
|
||||
'sid' : recovered_order.sid,
|
||||
'amount' : recovered_order.amount,
|
||||
'dt' : recovered_order.dt,
|
||||
@@ -4,15 +4,14 @@ import datetime
|
||||
import calendar
|
||||
import pytz
|
||||
import zipline.finance.risk as risk
|
||||
import zipline.test.factory as factory
|
||||
import zipline.util as qutil
|
||||
from zipline.utils import factory
|
||||
|
||||
from zipline.finance.trading import TradingEnvironment
|
||||
|
||||
class Risk(unittest.TestCase):
|
||||
|
||||
def setUp(self):
|
||||
qutil.configure_logging()
|
||||
#qutil.configure_logging()
|
||||
start_date = datetime.datetime(
|
||||
year=2006,
|
||||
month=1,
|
||||
@@ -354,4 +353,4 @@ RETURNS = [
|
||||
0.048 , -0.0307, -0.0357, 0.0033, -0.0412, -0.0407, 0.0455,
|
||||
0.0159, -0.0051, -0.0274, -0.0213, 0.0361, 0.0051, -0.0378,
|
||||
0.0084, 0.0066, -0.0103, -0.0037, 0.0478, -0.0278
|
||||
]
|
||||
]
|
||||
+5
-1
@@ -5,12 +5,16 @@ Zipline
|
||||
# This is *not* a place to dump arbitrary classes/modules for convenience,
|
||||
# it is a place to expose the public interfaces.
|
||||
|
||||
import protocol
|
||||
import protocol # namespace
|
||||
from core.monitor import Controller
|
||||
from lines import SimulatedTrading
|
||||
from core.host import ComponentHost
|
||||
from utils.protocol_utils import ndict
|
||||
|
||||
__all__ = [
|
||||
SimulatedTrading,
|
||||
Controller,
|
||||
ComponentHost,
|
||||
protocol,
|
||||
ndict
|
||||
]
|
||||
|
||||
@@ -1,562 +0,0 @@
|
||||
"""
|
||||
Commonly used messaging components.
|
||||
|
||||
Contains the base class for all components.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
import uuid
|
||||
import time
|
||||
import socket
|
||||
import gevent
|
||||
import traceback
|
||||
import humanhash
|
||||
|
||||
# pyzmq
|
||||
import zmq
|
||||
# gevent_zeromq
|
||||
import gevent_zeromq
|
||||
# zmq_ctypes
|
||||
#import zmq_ctypes
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
import zipline.util as qutil
|
||||
from zipline.gpoll import _Poller as GeventPoller
|
||||
from zipline.protocol import CONTROL_PROTOCOL, COMPONENT_STATE, \
|
||||
COMPONENT_FAILURE, BACKTEST_STATE, CONTROL_FRAME
|
||||
|
||||
|
||||
class Component(object):
|
||||
"""
|
||||
Base class for components. Defines the the base messaging
|
||||
interface for components.
|
||||
|
||||
:param addresses: a dict of name_string -> zmq port address strings.
|
||||
Must have the following entries
|
||||
|
||||
:param sync_address: socket address used for synchronizing the start of
|
||||
all workers, heartbeating, and exit notification
|
||||
will be used in REP/REQ sockets. Bind is always on
|
||||
the REP side.
|
||||
|
||||
:param data_address: socket address used for data sources to stream
|
||||
their records. Will be used in PUSH/PULL sockets
|
||||
between data sources and a Feed. Bind will always
|
||||
be on the PULL side (we always have N producers and
|
||||
1 consumer)
|
||||
|
||||
:param feed_address: socket address used to publish consolidated feed
|
||||
from serialization of data sources
|
||||
will be used in PUB/SUB sockets between Feed and
|
||||
Transforms. Bind is always on the PUB side.
|
||||
|
||||
:param merge_address: socket address used to publish transformed
|
||||
values. will be used in PUSH/PULL from many
|
||||
transforms to one Merge Bind will always be on
|
||||
the PULL side (we always have N producers and
|
||||
1 consumer)
|
||||
|
||||
:param result_address: socket address used to publish merged data
|
||||
source feed and transforms to clients will be
|
||||
used in PUB/SUB from one Merge to one or many
|
||||
clients. Bind is always on the PUB side.
|
||||
|
||||
bind/connect methods will return the correct socket type for each
|
||||
address.
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.zmq = None
|
||||
self.context = None
|
||||
self.addresses = None
|
||||
|
||||
self.out_socket = None
|
||||
self.killed = False
|
||||
self.controller = None
|
||||
# timeout after a full minute
|
||||
self.heartbeat_timeout = 60 *1000
|
||||
self.state_flag = COMPONENT_STATE.OK
|
||||
self.error_state = COMPONENT_FAILURE.NOFAILURE
|
||||
self.on_done = None
|
||||
|
||||
self._exception = None
|
||||
self.fail_time = None
|
||||
self.start_tic = None
|
||||
self.stop_tic = None
|
||||
self.note = None
|
||||
self.confirmed = False
|
||||
|
||||
# Humanhashes make this way easier to debug because they
|
||||
# stick in your mind unlike a 32 byte string of random hex.
|
||||
self.guid = uuid.uuid4()
|
||||
self.huid = humanhash.humanize(self.guid.hex)
|
||||
|
||||
self.init()
|
||||
|
||||
def init(self):
|
||||
"""
|
||||
Subclasses should override this to extend the setup for
|
||||
the class. Shouldn't have side effects.
|
||||
"""
|
||||
pass
|
||||
|
||||
# ------------
|
||||
# Core Methods
|
||||
# ------------
|
||||
|
||||
def open(self):
|
||||
"""
|
||||
Open the connections needed to start doing work.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def ready(self):
|
||||
"""
|
||||
Return ``True`` if and only if the component has finished execution.
|
||||
"""
|
||||
return self.state_flag in [COMPONENT_STATE.DONE, \
|
||||
COMPONENT_STATE.EXCEPTION]
|
||||
|
||||
def successful(self):
|
||||
"""
|
||||
Return ``True`` if and only if the component has finished execution
|
||||
successfully, that is, without raising an error.
|
||||
"""
|
||||
return self.state_flag == COMPONENT_STATE.DONE and not \
|
||||
self.exception
|
||||
|
||||
@property
|
||||
def exception(self):
|
||||
"""
|
||||
Holds the exception that the component failed on, or
|
||||
``None`` if the component has not failed.
|
||||
"""
|
||||
return self._exception
|
||||
|
||||
def do_work(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def init_zmq(self, flavor):
|
||||
"""
|
||||
ZMQ in all flavors. Have it your way.
|
||||
|
||||
mp - Distinct contexts | pyzmq
|
||||
thread - Same context | pyzmq
|
||||
green - Same context | gevent_zeromq
|
||||
pypy - Same context | zmq_ctypes
|
||||
|
||||
"""
|
||||
|
||||
if flavor == 'mp':
|
||||
self.zmq = zmq
|
||||
self.context = self.zmq.Context()
|
||||
self.zmq_poller = self.zmq.Poller
|
||||
return
|
||||
if flavor == 'thread':
|
||||
self.zmq = zmq
|
||||
self.context = self.zmq.Context.instance()
|
||||
self.zmq_poller = self.zmq.Poller
|
||||
return
|
||||
if flavor == 'green':
|
||||
self.zmq = gevent_zeromq.zmq
|
||||
self.context = self.zmq.Context.instance()
|
||||
self.zmq_poller = GeventPoller
|
||||
return
|
||||
if flavor == 'pypy':
|
||||
self.zmq = zmq
|
||||
self.context = self.zmq.Context.instance()
|
||||
self.zmq_poller = self.zmq.Poller
|
||||
return
|
||||
|
||||
raise Exception("Unknown ZeroMQ Flavor")
|
||||
|
||||
def _run(self):
|
||||
self.start_tic = time.time()
|
||||
|
||||
self.done = False # TODO: use state flag
|
||||
self.sockets = []
|
||||
|
||||
self.init_zmq(self.zmq_flavor)
|
||||
|
||||
self.setup_poller()
|
||||
|
||||
self.open()
|
||||
self.setup_sync()
|
||||
self.setup_control()
|
||||
|
||||
self.loop()
|
||||
self.shutdown()
|
||||
|
||||
self.stop_tic = time.time()
|
||||
|
||||
def run(self, catch_exceptions=True):
|
||||
"""
|
||||
Run the component.
|
||||
|
||||
Optionally takes an argument to catch and log all exceptions raised
|
||||
during execution ues this with care since it makes it very hard to
|
||||
debug since it mucks up your stacktraces.
|
||||
"""
|
||||
|
||||
if catch_exceptions:
|
||||
try:
|
||||
self._run()
|
||||
except Exception as exc:
|
||||
exc_info = sys.exc_info()
|
||||
self.signal_exception(exc)
|
||||
|
||||
# Reraise the exception
|
||||
raise exc_info[0], exc_info[1], exc_info[2]
|
||||
finally:
|
||||
|
||||
self.shutdown()
|
||||
self.teardown_sockets()
|
||||
|
||||
def working(self):
|
||||
"""
|
||||
Controls when the work loop will start and end
|
||||
|
||||
If we encounter an exception or signal done exit.
|
||||
|
||||
Overload for higher order behavior.
|
||||
"""
|
||||
return (not self.done)
|
||||
|
||||
def loop(self, lockstep=True):
|
||||
"""
|
||||
Loop to do work while we still have work to do.
|
||||
"""
|
||||
while self.working():
|
||||
self.confirm()
|
||||
self.do_work()
|
||||
|
||||
def confirm(self):
|
||||
"""
|
||||
Send a synchronization request to the host.
|
||||
"""
|
||||
if not self.confirmed:
|
||||
# TODO: proper framing
|
||||
self.sync_socket.send(self.get_id + ":RUN")
|
||||
|
||||
self.receive_sync_ack() # blocking
|
||||
self.confirmed = True
|
||||
|
||||
def runtime(self):
|
||||
if self.ready() and self.start_tic and self.stop_tic:
|
||||
return self.stop_tic - self.start_tic
|
||||
|
||||
# ----------------------------
|
||||
# Cleanup & Modes of Failure
|
||||
# ----------------------------
|
||||
|
||||
def teardown_sockets(self):
|
||||
"""
|
||||
Close all zmq sockets safely. This is universal, no matter
|
||||
where this is running it will need the sockets closed.
|
||||
"""
|
||||
#close all the sockets
|
||||
for sock in self.sockets:
|
||||
sock.close()
|
||||
|
||||
def shutdown(self):
|
||||
"""
|
||||
Clean shutdown.
|
||||
|
||||
Tear down after normal operation.
|
||||
"""
|
||||
if self.on_done:
|
||||
self.on_done()
|
||||
|
||||
def kill(self):
|
||||
"""
|
||||
Unclean shutdown.
|
||||
|
||||
Tear down ( fast ) as a mode of failure in the
|
||||
simulation or on service halt.
|
||||
|
||||
Context specific.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
# ----------------------
|
||||
# Internal Maintenance
|
||||
# ----------------------
|
||||
|
||||
def signal_exception(self, exc=None, scope=None):
|
||||
"""
|
||||
This is *very* important error tracking handler.
|
||||
|
||||
Will inform the system that the component has failed and
|
||||
how it has failed.
|
||||
"""
|
||||
|
||||
if scope == 'algo':
|
||||
self.error_state = COMPONENT_FAILURE.ALGOEXCEPT
|
||||
else:
|
||||
self.error_state = COMPONENT_FAILURE.HOSTEXCEPT
|
||||
|
||||
self.state_flag = COMPONENT_STATE.EXCEPTION
|
||||
# mark the time of failure so we can track the failure
|
||||
# progogation through the system.
|
||||
|
||||
self.stop_tic = time.time()
|
||||
|
||||
self._exception = exc
|
||||
exc_type, exc_value, exc_traceback = sys.exc_info()
|
||||
trace = '\n>>>'.join(traceback.format_exception(exc_type, exc_value, exc_traceback))
|
||||
|
||||
exception_frame = CONTROL_FRAME(
|
||||
CONTROL_PROTOCOL.EXCEPTION,
|
||||
trace
|
||||
)
|
||||
self.control_out.send(exception_frame)
|
||||
|
||||
qutil.LOGGER.exception("Unexpected error in run for {id}.".format(id=self.get_id))
|
||||
|
||||
def signal_done(self):
|
||||
"""
|
||||
Notify down stream components that we're done.
|
||||
"""
|
||||
|
||||
self.state_flag = COMPONENT_STATE.DONE
|
||||
|
||||
if self.out_socket:
|
||||
self.out_socket.send(str(CONTROL_PROTOCOL.DONE))
|
||||
|
||||
#notify host we're done
|
||||
# TODO: proper framing
|
||||
self.sync_socket.send(self.get_id + ":" + str(CONTROL_PROTOCOL.DONE))
|
||||
|
||||
#notify controller we're done
|
||||
done_frame = CONTROL_FRAME(
|
||||
CONTROL_PROTOCOL.DONE,
|
||||
''
|
||||
)
|
||||
self.control_out.send(done_frame)
|
||||
|
||||
self.receive_sync_ack()
|
||||
#notify internal work look that we're done
|
||||
self.done = True # TODO: use state flag
|
||||
|
||||
qutil.LOGGER.info("[%s] DONE" % self.get_id)
|
||||
|
||||
# -----------
|
||||
# Messaging
|
||||
# -----------
|
||||
|
||||
def setup_poller(self):
|
||||
"""
|
||||
Setup the poller used for multiplexing the incoming data
|
||||
handling sockets.
|
||||
"""
|
||||
|
||||
# Initializes the poller class specified by the flavor of
|
||||
# ZeroMQ. Either zmq.Poller or gpoll.Poller .
|
||||
self.poll = self.zmq_poller()
|
||||
|
||||
def receive_sync_ack(self):
|
||||
"""
|
||||
Wait for synchronization reply from the host.
|
||||
|
||||
DEPRECATED, left in for compatability for now.
|
||||
"""
|
||||
|
||||
socks = dict(self.sync_poller.poll(self.heartbeat_timeout))
|
||||
if self.sync_socket in socks and socks[self.sync_socket] == self.zmq.POLLIN:
|
||||
message = self.sync_socket.recv()
|
||||
#else:
|
||||
#raise Exception("Sync ack timed out on response for {id}".format(id=self.get_id))
|
||||
|
||||
def bind_data(self):
|
||||
return self.bind_pull_socket(self.addresses['data_address'])
|
||||
|
||||
def connect_data(self):
|
||||
return self.connect_push_socket(self.addresses['data_address'])
|
||||
|
||||
def bind_feed(self):
|
||||
return self.bind_pub_socket(self.addresses['feed_address'])
|
||||
|
||||
def connect_feed(self):
|
||||
return self.connect_sub_socket(self.addresses['feed_address'])
|
||||
|
||||
def bind_merge(self):
|
||||
return self.bind_pull_socket(self.addresses['merge_address'])
|
||||
|
||||
def connect_merge(self):
|
||||
return self.connect_push_socket(self.addresses['merge_address'])
|
||||
|
||||
def bind_result(self):
|
||||
return self.bind_pub_socket(self.addresses['result_address'])
|
||||
|
||||
def connect_result(self):
|
||||
return self.connect_sub_socket(self.addresses['result_address'])
|
||||
|
||||
def bind_pull_socket(self, addr):
|
||||
pull_socket = self.context.socket(self.zmq.PULL)
|
||||
pull_socket.bind(addr)
|
||||
self.poll.register(pull_socket, self.zmq.POLLIN)
|
||||
|
||||
self.sockets.append(pull_socket)
|
||||
|
||||
return pull_socket
|
||||
|
||||
def connect_push_socket(self, addr):
|
||||
push_socket = self.context.socket(self.zmq.PUSH)
|
||||
push_socket.connect(addr)
|
||||
#push_socket.setsockopt(self.zmq.LINGER,0)
|
||||
self.sockets.append(push_socket)
|
||||
self.out_socket = push_socket
|
||||
|
||||
return push_socket
|
||||
|
||||
def bind_pub_socket(self, addr):
|
||||
pub_socket = self.context.socket(self.zmq.PUB)
|
||||
pub_socket.bind(addr)
|
||||
#pub_socket.setsockopt(self.zmq.LINGER,0)
|
||||
self.out_socket = pub_socket
|
||||
|
||||
return pub_socket
|
||||
|
||||
def connect_sub_socket(self, addr):
|
||||
sub_socket = self.context.socket(self.zmq.SUB)
|
||||
sub_socket.connect(addr)
|
||||
sub_socket.setsockopt(self.zmq.SUBSCRIBE,'')
|
||||
self.sockets.append(sub_socket)
|
||||
|
||||
self.poll.register(sub_socket, self.zmq.POLLIN)
|
||||
|
||||
return sub_socket
|
||||
|
||||
def setup_control(self):
|
||||
"""
|
||||
Set up the control socket. Used to monitor the
|
||||
overall status of the simulation and to forcefully tear
|
||||
down the simulation in case of a failure.
|
||||
"""
|
||||
|
||||
# Allow for the possibility of not having a controller,
|
||||
# possibly the zipline devsimulator may not want this.
|
||||
if not self.controller:
|
||||
return
|
||||
|
||||
self.control_out = self.controller.message_sender(
|
||||
identity = self.get_id,
|
||||
context = self.context,
|
||||
)
|
||||
|
||||
self.control_in = self.controller.message_listener(
|
||||
context = self.context
|
||||
)
|
||||
|
||||
self.poll.register(self.control_in, self.zmq.POLLIN)
|
||||
self.sockets.extend([self.control_in, self.control_out])
|
||||
|
||||
def setup_sync(self):
|
||||
"""
|
||||
Setup the sync socket and poller. ( Connect )
|
||||
|
||||
DEPRECATED, left in for compatability for now.
|
||||
"""
|
||||
|
||||
qutil.LOGGER.debug("Connecting sync client for {id}".format(id=self.get_id))
|
||||
|
||||
self.sync_socket = self.context.socket(self.zmq.REQ)
|
||||
self.sync_socket.connect(self.addresses['sync_address'])
|
||||
#self.sync_socket.setsockopt(self.zmq.LINGER,0)
|
||||
|
||||
self.sync_poller = self.zmq_poller()
|
||||
self.sync_poller.register(self.sync_socket, self.zmq.POLLIN)
|
||||
|
||||
self.sockets.append(self.sync_socket)
|
||||
|
||||
# ---------------------
|
||||
# Description and Debug
|
||||
# ---------------------
|
||||
|
||||
def extern_logger(self):
|
||||
"""
|
||||
Pipe logs out to a provided logging interface.
|
||||
"""
|
||||
pass
|
||||
|
||||
def setup_extern_logger(self):
|
||||
"""
|
||||
Pipe logs out to a provided logging interface.
|
||||
"""
|
||||
pass
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
"""
|
||||
The descriptive name of the component.
|
||||
"""
|
||||
# Prevents the bug that Thomas ran into
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
"""
|
||||
The data flow type of the component.
|
||||
|
||||
- ``SOURCE``
|
||||
- ``CONDUIT``
|
||||
- ``SINK``
|
||||
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@property
|
||||
def get_pure(self):
|
||||
"""
|
||||
Describes whehter this component purely functional,
|
||||
i.e. for a given set of inputs is it guaranteed to
|
||||
always give the same output . Components that are
|
||||
side-effectful are, generally, not pure.
|
||||
"""
|
||||
return False
|
||||
|
||||
def note(self):
|
||||
"""
|
||||
Information about the component. Mostly used for testing.
|
||||
"""
|
||||
|
||||
def get_note(self):
|
||||
return self.note or ''
|
||||
|
||||
def debug(self):
|
||||
"""
|
||||
Debug information about the component.
|
||||
"""
|
||||
return {
|
||||
'id' : self.get_id ,
|
||||
'huid' : self.huid ,
|
||||
'host' : socket.gethostname() ,
|
||||
'pid' : os.getpid() ,
|
||||
'memaddress' : hex(id(self)) ,
|
||||
'ready' : self.successful() ,
|
||||
'succesfull' : self.ready() ,
|
||||
}
|
||||
|
||||
def __len__(self):
|
||||
"""
|
||||
Some components overload this for debug purposes
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def __repr__(self):
|
||||
"""
|
||||
Return a usefull string representation of the component
|
||||
to indicate its type, unique identifier, and computational
|
||||
context identifier name.
|
||||
"""
|
||||
|
||||
return "<{name} {uuid} at {host} {pid} {pointer}>".format(
|
||||
name = self.get_id ,
|
||||
uuid = self.huid ,
|
||||
host = socket.gethostname() ,
|
||||
pid = os.getpid() ,
|
||||
pointer = hex(id(self)) ,
|
||||
)
|
||||
@@ -0,0 +1,11 @@
|
||||
from feed import Feed
|
||||
from merge import Merge
|
||||
from passthrough import PassthroughTransform
|
||||
from datasource import DataSource
|
||||
|
||||
__all__ = [
|
||||
Feed,
|
||||
Merge,
|
||||
PassthroughTransform,
|
||||
DataSource,
|
||||
]
|
||||
@@ -0,0 +1,66 @@
|
||||
"""
|
||||
Commonly used messaging components.
|
||||
"""
|
||||
|
||||
import logging
|
||||
|
||||
import zipline.protocol as zp
|
||||
from zipline.core.component import Component
|
||||
from zipline.protocol import COMPONENT_TYPE
|
||||
|
||||
LOGGER = logging.getLogger('ZiplineLogger')
|
||||
|
||||
class DataSource(Component):
|
||||
"""
|
||||
Baseclass for data sources. Subclass and implement send_all - usually this
|
||||
means looping through all records in a store, converting to a dict, and
|
||||
calling send(map).
|
||||
|
||||
Every datasource has a dict property to hold filters::
|
||||
- key -- name of the filter, e.g. SID
|
||||
- value -- a primitive representing the filter. e.g. a list of ints.
|
||||
|
||||
Modify the datasource's filters via the set_filter(name, value)
|
||||
"""
|
||||
def __init__(self, source_id):
|
||||
Component.__init__(self)
|
||||
|
||||
self.id = source_id
|
||||
self.init()
|
||||
self.filter = {}
|
||||
|
||||
def init(self):
|
||||
self.cur_event = None
|
||||
|
||||
def set_filter(self, name, value):
|
||||
self.filter[name] = value
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return self.id
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
return COMPONENT_TYPE.SOURCE
|
||||
|
||||
def open(self):
|
||||
self.data_socket = self.connect_data()
|
||||
|
||||
def send(self, event):
|
||||
"""
|
||||
Emit data.
|
||||
"""
|
||||
assert isinstance(event, zp.ndict)
|
||||
|
||||
event['source_id'] = self.get_id
|
||||
event['type'] = self.get_type
|
||||
|
||||
try:
|
||||
ds_frame = self.frame(event)
|
||||
except zp.INVALID_DATASOURCE_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
self.data_socket.send(ds_frame)
|
||||
|
||||
def frame(self, event):
|
||||
return zp.DATASOURCE_FRAME(event)
|
||||
@@ -0,0 +1,209 @@
|
||||
import logging
|
||||
from collections import Counter
|
||||
|
||||
from zipline.core.component import Component
|
||||
import zipline.protocol as zp
|
||||
|
||||
from zipline.protocol import CONTROL_PROTOCOL, COMPONENT_TYPE, \
|
||||
CONTROL_FRAME, CONTROL_UNFRAME
|
||||
|
||||
LOGGER = logging.getLogger('ZiplineLogger')
|
||||
|
||||
class Feed(Component):
|
||||
"""
|
||||
Connects to N PULL sockets, publishing all messages received to a PUB
|
||||
socket. Published messages are guaranteed to be in chronological order
|
||||
based on message property dt. Expects to be instantiated in one execution
|
||||
context (thread, process, etc) and run in another.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
Component.__init__(self)
|
||||
|
||||
self.sent_count = 0
|
||||
self.received_count = 0
|
||||
self.draining = False
|
||||
self.ds_finished_counter = 0
|
||||
|
||||
# Depending on the size of this, might want to use a data
|
||||
# structure with better asymptotics.
|
||||
self.data_buffer = {}
|
||||
|
||||
# source_id -> integer count
|
||||
self.sent_counters = Counter()
|
||||
self.recv_counters = Counter()
|
||||
|
||||
def init(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return "FEED"
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
return COMPONENT_TYPE.CONDUIT
|
||||
|
||||
# -------------
|
||||
# Core Methods
|
||||
# -------------
|
||||
|
||||
def open(self):
|
||||
self.pull_socket = self.bind_data()
|
||||
self.feed_socket = self.bind_feed()
|
||||
|
||||
def do_work(self):
|
||||
# wait for synchronization reply from the host
|
||||
socks = dict(self.poll.poll(self.heartbeat_timeout))
|
||||
|
||||
# TODO: Abstract this out, maybe on base component
|
||||
if self.control_in in socks and socks[self.control_in] == self.zmq.POLLIN:
|
||||
msg = self.control_in.recv()
|
||||
event, payload = CONTROL_UNFRAME(msg)
|
||||
|
||||
# -- Heartbeat --
|
||||
if event == CONTROL_PROTOCOL.HEARTBEAT:
|
||||
# Heart outgoing
|
||||
heartbeat_frame = CONTROL_FRAME(
|
||||
CONTROL_PROTOCOL.OK,
|
||||
payload
|
||||
)
|
||||
self.control_out.send(heartbeat_frame)
|
||||
|
||||
# -- Soft Kill --
|
||||
elif event == CONTROL_PROTOCOL.SHUTDOWN:
|
||||
self.signal_done()
|
||||
self.shutdown()
|
||||
|
||||
# -- Hard Kill --
|
||||
elif event == CONTROL_PROTOCOL.KILL:
|
||||
self.kill()
|
||||
|
||||
|
||||
if self.pull_socket in socks and socks[self.pull_socket] == self.zmq.POLLIN:
|
||||
message = self.pull_socket.recv()
|
||||
|
||||
if message == str(CONTROL_PROTOCOL.DONE):
|
||||
self.ds_finished_counter += 1
|
||||
|
||||
if len(self.data_buffer) == self.ds_finished_counter:
|
||||
#drain any remaining messages in the buffer
|
||||
LOGGER.debug("draining feed")
|
||||
self.drain()
|
||||
self.signal_done()
|
||||
else:
|
||||
try:
|
||||
event = self.unframe(message)
|
||||
# deserialization error
|
||||
except zp.INVALID_DATASOURCE_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
try:
|
||||
self.append(event)
|
||||
self.send_next()
|
||||
|
||||
# Invalid message
|
||||
except zp.INVALID_DATASOURCE_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
def unframe(self, msg):
|
||||
return zp.DATASOURCE_UNFRAME(msg)
|
||||
|
||||
def frame(self, event):
|
||||
return zp.FEED_FRAME(event)
|
||||
|
||||
# -------------
|
||||
# Flow Control
|
||||
# -------------
|
||||
|
||||
def drain(self):
|
||||
"""
|
||||
Send all messages in the buffer.
|
||||
"""
|
||||
self.draining = True
|
||||
while self.pending_messages() > 0:
|
||||
self.send_next()
|
||||
|
||||
def send_next(self):
|
||||
"""
|
||||
Send the (chronologically) next message in the buffer.
|
||||
"""
|
||||
if not (self.is_full() or self.draining):
|
||||
return
|
||||
|
||||
event = self.next()
|
||||
if(event != None):
|
||||
self.feed_socket.send(self.frame(event), self.zmq.NOBLOCK)
|
||||
self.sent_counters[event.source_id] += 1
|
||||
self.sent_count += 1
|
||||
|
||||
def append(self, event):
|
||||
"""
|
||||
Add an event to the buffer for the source specified by
|
||||
source_id.
|
||||
"""
|
||||
self.data_buffer[event.source_id].append(event)
|
||||
self.recv_counters[event.source_id] += 1
|
||||
self.received_count += 1
|
||||
|
||||
def next(self):
|
||||
"""
|
||||
Get the next message in chronological order.
|
||||
"""
|
||||
if not(self.is_full() or self.draining):
|
||||
return
|
||||
|
||||
cur_source = None
|
||||
earliest_source = None
|
||||
earliest_event = None
|
||||
#iterate over the queues of events from all sources
|
||||
#(1 queue per datasource)
|
||||
for events in self.data_buffer.values():
|
||||
if len(events) == 0:
|
||||
continue
|
||||
cur_source = events
|
||||
first_in_list = events[0]
|
||||
if first_in_list.dt == None:
|
||||
#this is a filler event, discard
|
||||
events.pop(0)
|
||||
continue
|
||||
|
||||
if (earliest_event == None) or (first_in_list.dt <= earliest_event.dt):
|
||||
earliest_event = first_in_list
|
||||
earliest_source = cur_source
|
||||
|
||||
if earliest_event != None:
|
||||
return earliest_source.pop(0)
|
||||
|
||||
def is_full(self):
|
||||
"""
|
||||
Indicates whether the buffer has messages in buffer for
|
||||
all un-DONE, blocking sources.
|
||||
"""
|
||||
for source_id, events in self.data_buffer.iteritems():
|
||||
if len(events) == 0:
|
||||
return False
|
||||
return True
|
||||
|
||||
def pending_messages(self):
|
||||
"""
|
||||
Returns the count of all events from all sources in the
|
||||
buffer.
|
||||
"""
|
||||
total = 0
|
||||
for events in self.data_buffer.values():
|
||||
total += len(events)
|
||||
return total
|
||||
|
||||
def add_source(self, source_id):
|
||||
"""
|
||||
Add a data source to the buffer.
|
||||
"""
|
||||
self.data_buffer[source_id] = []
|
||||
|
||||
def __len__(self):
|
||||
"""
|
||||
Buffer's length is same as internal map holding separate
|
||||
sorted arrays of events keyed by source id.
|
||||
"""
|
||||
return len(self.data_buffer)
|
||||
@@ -0,0 +1,68 @@
|
||||
from feed import Feed
|
||||
|
||||
import zipline.protocol as zp
|
||||
from zipline.protocol import COMPONENT_TYPE
|
||||
|
||||
# TODO: By Liskov merge must *be* a feed, don't believe this is
|
||||
# the case.
|
||||
|
||||
class Merge(Feed):
|
||||
"""
|
||||
Merges multiple streams of events into single messages.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
Feed.__init__(self)
|
||||
|
||||
self.init()
|
||||
|
||||
def init(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return "MERGE"
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
return COMPONENT_TYPE.CONDUIT
|
||||
|
||||
def open(self):
|
||||
self.pull_socket = self.bind_merge()
|
||||
self.feed_socket = self.bind_result()
|
||||
|
||||
def next(self):
|
||||
"""Get the next merged message from the feed buffer."""
|
||||
if not (self.is_full() or self.draining):
|
||||
return
|
||||
|
||||
if self.pending_messages() == 0:
|
||||
return
|
||||
|
||||
#get the raw event from the passthrough transform.
|
||||
result = self.data_buffer[zp.TRANSFORM_TYPE.PASSTHROUGH].pop(0).PASSTHROUGH
|
||||
for source, events in self.data_buffer.iteritems():
|
||||
if source == zp.TRANSFORM_TYPE.PASSTHROUGH:
|
||||
continue
|
||||
if len(events) > 0:
|
||||
cur = events.pop(0)
|
||||
result.merge(cur)
|
||||
return result
|
||||
|
||||
def unframe(self, msg):
|
||||
return zp.TRANSFORM_UNFRAME(msg)
|
||||
|
||||
def frame(self, event):
|
||||
return zp.MERGE_FRAME(event)
|
||||
|
||||
def append(self, event):
|
||||
"""
|
||||
:param event: a ndict with one entry. key is the name of the
|
||||
transform, value is the transformed value.
|
||||
Add an event to the buffer for the source specified by
|
||||
source_id.
|
||||
"""
|
||||
|
||||
self.data_buffer[event.keys()[0]].append(event)
|
||||
self.received_count += 1
|
||||
|
||||
@@ -0,0 +1,35 @@
|
||||
import zipline.protocol as zp
|
||||
from zipline.transforms import BaseTransform
|
||||
|
||||
from zipline.protocol import CONTROL_PROTOCOL, COMPONENT_TYPE, \
|
||||
COMPONENT_STATE, CONTROL_FRAME, CONTROL_UNFRAME
|
||||
|
||||
class PassthroughTransform(BaseTransform):
|
||||
"""
|
||||
A bypass transform which is also an identity transform::
|
||||
|
||||
+-------+
|
||||
+---| f |--->
|
||||
+-------+
|
||||
+------id------->
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
BaseTransform.__init__(self, "PASSTHROUGH")
|
||||
self.init(**kwargs)
|
||||
|
||||
def init(self, **kwargs):
|
||||
pass
|
||||
|
||||
@property
|
||||
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 {
|
||||
'name' : zp.TRANSFORM_TYPE.PASSTHROUGH,
|
||||
'value' : zp.FEED_FRAME(event)
|
||||
}
|
||||
@@ -0,0 +1,9 @@
|
||||
from host import ComponentHost
|
||||
from component import Component
|
||||
from monitor import Controller
|
||||
|
||||
__all__ = [
|
||||
Component,
|
||||
Controller,
|
||||
ComponentHost
|
||||
]
|
||||
@@ -1,6 +1,4 @@
|
||||
"""
|
||||
Commonly used messaging components.
|
||||
|
||||
Contains the base class for all components.
|
||||
"""
|
||||
|
||||
@@ -9,7 +7,7 @@ import sys
|
||||
import uuid
|
||||
import time
|
||||
import socket
|
||||
import gevent
|
||||
import logging
|
||||
import traceback
|
||||
import humanhash
|
||||
|
||||
@@ -20,13 +18,11 @@ import gevent_zeromq
|
||||
# zmq_ctypes
|
||||
#import zmq_ctypes
|
||||
|
||||
from datetime import datetime
|
||||
|
||||
import zipline.util as qutil
|
||||
from zipline.gpoll import _Poller as GeventPoller
|
||||
from zipline.utils.gpoll import _Poller as GeventPoller
|
||||
from zipline.protocol import CONTROL_PROTOCOL, COMPONENT_STATE, \
|
||||
COMPONENT_FAILURE, BACKTEST_STATE, CONTROL_FRAME
|
||||
COMPONENT_FAILURE, CONTROL_FRAME
|
||||
|
||||
LOGGER = logging.getLogger('ZiplineLogger')
|
||||
|
||||
class Component(object):
|
||||
"""
|
||||
@@ -243,7 +239,7 @@ class Component(object):
|
||||
|
||||
self.receive_sync_ack() # blocking
|
||||
self.confirmed = True
|
||||
|
||||
|
||||
def runtime(self):
|
||||
if self.ready() and self.start_tic and self.stop_tic:
|
||||
return self.stop_tic - self.start_tic
|
||||
@@ -314,7 +310,7 @@ class Component(object):
|
||||
)
|
||||
self.control_out.send(exception_frame)
|
||||
|
||||
qutil.LOGGER.exception("Unexpected error in run for {id}.".format(id=self.get_id))
|
||||
LOGGER.exception("Unexpected error in run for {id}.".format(id=self.get_id))
|
||||
|
||||
def signal_done(self):
|
||||
"""
|
||||
@@ -341,7 +337,7 @@ class Component(object):
|
||||
#notify internal work look that we're done
|
||||
self.done = True # TODO: use state flag
|
||||
|
||||
qutil.LOGGER.info("[%s] DONE" % self.get_id)
|
||||
LOGGER.info("[%s] DONE" % self.get_id)
|
||||
|
||||
# -----------
|
||||
# Messaging
|
||||
@@ -461,7 +457,7 @@ class Component(object):
|
||||
DEPRECATED, left in for compatability for now.
|
||||
"""
|
||||
|
||||
qutil.LOGGER.debug("Connecting sync client for {id}".format(id=self.get_id))
|
||||
LOGGER.debug("Connecting sync client for {id}".format(id=self.get_id))
|
||||
|
||||
self.sync_socket = self.context.socket(self.zmq.REQ)
|
||||
self.sync_socket.connect(self.addresses['sync_address'])
|
||||
|
||||
@@ -3,11 +3,7 @@ Simulator hosts all the components necessary to execute a simluation. See :py:me
|
||||
"""
|
||||
|
||||
import threading
|
||||
import mock
|
||||
from collections import defaultdict
|
||||
from zipline.monitor import Controller
|
||||
from zipline.messaging import ComponentHost
|
||||
import zipline.util as qutil
|
||||
from zipline.core import ComponentHost
|
||||
|
||||
class AddressAllocator(object):
|
||||
|
||||
@@ -35,7 +31,7 @@ class Simulator(ComponentHost):
|
||||
ComponentHost.__init__(self, addresses)
|
||||
self.subthreads = []
|
||||
self.running = False
|
||||
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return 'Simple Simulator'
|
||||
@@ -0,0 +1,164 @@
|
||||
import logging
|
||||
import datetime
|
||||
|
||||
from component import Component
|
||||
|
||||
from zipline.transforms import BaseTransform
|
||||
from zipline.components import Feed, Merge, PassthroughTransform, \
|
||||
DataSource
|
||||
from zipline.protocol import CONTROL_PROTOCOL, COMPONENT_STATE
|
||||
|
||||
LOGGER = logging.getLogger('ZiplineLogger')
|
||||
|
||||
class ComponentHost(Component):
|
||||
"""
|
||||
Components that can launch multiple sub-components, synchronize their
|
||||
start, and then wait for all components to be finished.
|
||||
"""
|
||||
|
||||
def __init__(self, addresses):
|
||||
Component.__init__(self)
|
||||
self.addresses = addresses
|
||||
self.running = False
|
||||
|
||||
self.init()
|
||||
|
||||
def init(self):
|
||||
assert hasattr(self, 'zmq_flavor'), """
|
||||
You must specify a flavor of ZeroMQ for all
|
||||
ComponentHost subclasses. """
|
||||
|
||||
# Component Registry, keyed by get_id
|
||||
# ----------------------
|
||||
self.components = {}
|
||||
# ----------------------
|
||||
# Internal Registry, keyed by guid
|
||||
self._components = {}
|
||||
# ----------------------
|
||||
|
||||
self.sync_register = {}
|
||||
self.timeout = datetime.timedelta(seconds=60)
|
||||
|
||||
self.feed = Feed()
|
||||
self.merge = Merge()
|
||||
self.passthrough = PassthroughTransform()
|
||||
self.controller = None
|
||||
|
||||
#register the feed and the merge
|
||||
self.register_components([self.feed, self.merge, self.passthrough])
|
||||
|
||||
def register_controller(self, controller):
|
||||
"""
|
||||
Add the given components to the registry. Establish
|
||||
communication with them.
|
||||
"""
|
||||
if self.controller != None:
|
||||
raise Exception("There can be only one!")
|
||||
|
||||
self.controller = controller
|
||||
self.controller.zmq_flavor = self.zmq_flavor
|
||||
|
||||
# Propogate the controller to all the subcomponents
|
||||
for component in self.components.itervalues():
|
||||
component.controller = controller
|
||||
|
||||
def register_components(self, components):
|
||||
"""
|
||||
Add the given components to the registry. Establish
|
||||
communication with them.
|
||||
"""
|
||||
assert isinstance(components, list)
|
||||
for component in components:
|
||||
|
||||
component.addresses = self.addresses
|
||||
component.controller = self.controller
|
||||
|
||||
# Hosts share their zmq flavor with hosted components
|
||||
component.zmq_flavor = self.zmq_flavor
|
||||
|
||||
self._components[component.guid] = component
|
||||
self.components[component.get_id] = component
|
||||
self.sync_register[component.get_id] = datetime.datetime.utcnow()
|
||||
|
||||
if isinstance(component, DataSource):
|
||||
self.feed.add_source(component.get_id)
|
||||
if isinstance(component, BaseTransform):
|
||||
self.merge.add_source(component.get_id)
|
||||
|
||||
def unregister_component(self, component_id):
|
||||
del self.components[component_id]
|
||||
del self.sync_register[component_id]
|
||||
|
||||
def setup_sync(self):
|
||||
"""
|
||||
Setup the sync socket and poller. ( Bind )
|
||||
"""
|
||||
LOGGER.debug("Connecting sync server.")
|
||||
|
||||
self.sync_socket = self.context.socket(self.zmq.REP)
|
||||
self.sync_socket.bind(self.addresses['sync_address'])
|
||||
|
||||
self.sync_poller = self.zmq_poller()
|
||||
self.sync_poller.register(self.sync_socket, self.zmq.POLLIN)
|
||||
|
||||
self.sockets.append(self.sync_socket)
|
||||
|
||||
def open(self):
|
||||
for component in self.components.values():
|
||||
self.launch_component(component)
|
||||
self.launch_controller()
|
||||
|
||||
def is_running(self):
|
||||
"""
|
||||
DEPRECATED, left in for compatability for now.
|
||||
"""
|
||||
|
||||
cur_time = datetime.datetime.utcnow()
|
||||
|
||||
if len(self.components) == 0:
|
||||
LOGGER.info("Component register is empty.")
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def loop(self, lockstep=True):
|
||||
|
||||
while self.is_running():
|
||||
# wait for synchronization request at start, and DONE at end.
|
||||
# don't timeout.
|
||||
socks = dict(self.sync_poller.poll())
|
||||
|
||||
if self.sync_socket in socks and socks[self.sync_socket] == self.zmq.POLLIN:
|
||||
msg = self.sync_socket.recv()
|
||||
|
||||
try:
|
||||
parts = msg.split(':')
|
||||
sync_id, status = parts
|
||||
except ValueError as exc:
|
||||
self.signal_exception(exc)
|
||||
|
||||
if status == str(CONTROL_PROTOCOL.DONE): # TODO: other way around
|
||||
LOGGER.debug("{id} is DONE".format(id=sync_id))
|
||||
self.unregister_component(sync_id)
|
||||
self.state_flag = COMPONENT_STATE.DONE
|
||||
else:
|
||||
self.sync_register[sync_id] = datetime.datetime.utcnow()
|
||||
|
||||
#qutil.LOGGER.info("confirmed {id}".format(id=msg))
|
||||
# send synchronization reply
|
||||
self.sync_socket.send('ack', self.zmq.NOBLOCK)
|
||||
|
||||
# ------------------
|
||||
# Simulation Control
|
||||
# ------------------
|
||||
|
||||
def launch_controller(self, controller):
|
||||
raise NotImplementedError
|
||||
|
||||
def launch_component(self, component):
|
||||
raise NotImplementedError
|
||||
|
||||
def teardown_component(self, component):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
@@ -1,636 +0,0 @@
|
||||
"""
|
||||
Commonly used messaging components.
|
||||
"""
|
||||
|
||||
import datetime
|
||||
|
||||
from collections import Counter
|
||||
|
||||
import zipline.util as qutil
|
||||
from zipline.component import Component
|
||||
import zipline.protocol as zp
|
||||
from zipline.protocol import CONTROL_PROTOCOL, COMPONENT_TYPE, \
|
||||
COMPONENT_STATE, CONTROL_FRAME, CONTROL_UNFRAME
|
||||
|
||||
class ComponentHost(Component):
|
||||
"""
|
||||
Components that can launch multiple sub-components, synchronize their
|
||||
start, and then wait for all components to be finished.
|
||||
"""
|
||||
|
||||
def __init__(self, addresses):
|
||||
Component.__init__(self)
|
||||
self.addresses = addresses
|
||||
self.running = False
|
||||
|
||||
self.init()
|
||||
|
||||
def init(self):
|
||||
assert hasattr(self, 'zmq_flavor'), \
|
||||
""" You must specify a flavor of ZeroMQ for all
|
||||
ComponentHost subclasses. """
|
||||
|
||||
# Component Registry, keyed by get_id
|
||||
# ----------------------
|
||||
self.components = {}
|
||||
# ----------------------
|
||||
# Internal Registry, keyed by guid
|
||||
self._components = {}
|
||||
# ----------------------
|
||||
|
||||
self.sync_register = {}
|
||||
self.timeout = datetime.timedelta(seconds=60)
|
||||
|
||||
self.feed = Feed()
|
||||
self.merge = Merge()
|
||||
self.passthrough = PassthroughTransform()
|
||||
self.controller = None
|
||||
|
||||
#register the feed and the merge
|
||||
self.register_components([self.feed, self.merge, self.passthrough])
|
||||
|
||||
def register_controller(self, controller):
|
||||
"""
|
||||
Add the given components to the registry. Establish
|
||||
communication with them.
|
||||
"""
|
||||
if self.controller != None:
|
||||
raise Exception("There can be only one!")
|
||||
|
||||
self.controller = controller
|
||||
self.controller.zmq_flavor = self.zmq_flavor
|
||||
|
||||
# Propogate the controller to all the subcomponents
|
||||
for component in self.components.itervalues():
|
||||
component.controller = controller
|
||||
|
||||
def register_components(self, components):
|
||||
"""
|
||||
Add the given components to the registry. Establish
|
||||
communication with them.
|
||||
"""
|
||||
assert isinstance(components, list)
|
||||
for component in components:
|
||||
|
||||
component.addresses = self.addresses
|
||||
component.controller = self.controller
|
||||
|
||||
# Hosts share their zmq flavor with hosted components
|
||||
component.zmq_flavor = self.zmq_flavor
|
||||
|
||||
self._components[component.guid] = component
|
||||
self.components[component.get_id] = component
|
||||
self.sync_register[component.get_id] = datetime.datetime.utcnow()
|
||||
|
||||
if isinstance(component, DataSource):
|
||||
self.feed.add_source(component.get_id)
|
||||
if isinstance(component, BaseTransform):
|
||||
self.merge.add_source(component.get_id)
|
||||
|
||||
def unregister_component(self, component_id):
|
||||
del self.components[component_id]
|
||||
del self.sync_register[component_id]
|
||||
|
||||
def setup_sync(self):
|
||||
"""
|
||||
Setup the sync socket and poller. ( Bind )
|
||||
"""
|
||||
qutil.LOGGER.debug("Connecting sync server.")
|
||||
|
||||
self.sync_socket = self.context.socket(self.zmq.REP)
|
||||
self.sync_socket.bind(self.addresses['sync_address'])
|
||||
|
||||
self.sync_poller = self.zmq_poller()
|
||||
self.sync_poller.register(self.sync_socket, self.zmq.POLLIN)
|
||||
|
||||
self.sockets.append(self.sync_socket)
|
||||
|
||||
def open(self):
|
||||
for component in self.components.values():
|
||||
self.launch_component(component)
|
||||
self.launch_controller()
|
||||
|
||||
def is_running(self):
|
||||
"""
|
||||
DEPRECATED, left in for compatability for now.
|
||||
"""
|
||||
|
||||
cur_time = datetime.datetime.utcnow()
|
||||
|
||||
if len(self.components) == 0:
|
||||
qutil.LOGGER.info("Component register is empty.")
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def loop(self, lockstep=True):
|
||||
|
||||
while self.is_running():
|
||||
# wait for synchronization request at start, and DONE at end.
|
||||
# don't timeout.
|
||||
socks = dict(self.sync_poller.poll())
|
||||
|
||||
if self.sync_socket in socks and socks[self.sync_socket] == self.zmq.POLLIN:
|
||||
msg = self.sync_socket.recv()
|
||||
|
||||
try:
|
||||
parts = msg.split(':')
|
||||
sync_id, status = parts
|
||||
except ValueError as exc:
|
||||
self.signal_exception(exc)
|
||||
|
||||
if status == str(CONTROL_PROTOCOL.DONE): # TODO: other way around
|
||||
#qutil.LOGGER.debug("{id} is DONE".format(id=sync_id))
|
||||
self.unregister_component(sync_id)
|
||||
self.state_flag = COMPONENT_STATE.DONE
|
||||
else:
|
||||
self.sync_register[sync_id] = datetime.datetime.utcnow()
|
||||
|
||||
#qutil.LOGGER.info("confirmed {id}".format(id=msg))
|
||||
# send synchronization reply
|
||||
self.sync_socket.send('ack', self.zmq.NOBLOCK)
|
||||
|
||||
# ------------------
|
||||
# Simulation Control
|
||||
# ------------------
|
||||
|
||||
def launch_controller(self, controller):
|
||||
raise NotImplementedError
|
||||
|
||||
def launch_component(self, component):
|
||||
raise NotImplementedError
|
||||
|
||||
def teardown_component(self, component):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class Feed(Component):
|
||||
"""
|
||||
Connects to N PULL sockets, publishing all messages received to a PUB
|
||||
socket. Published messages are guaranteed to be in chronological order
|
||||
based on message property dt. Expects to be instantiated in one execution
|
||||
context (thread, process, etc) and run in another.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
Component.__init__(self)
|
||||
|
||||
self.sent_count = 0
|
||||
self.received_count = 0
|
||||
self.draining = False
|
||||
self.ds_finished_counter = 0
|
||||
|
||||
# Depending on the size of this, might want to use a data
|
||||
# structure with better asymptotics.
|
||||
self.data_buffer = {}
|
||||
|
||||
# source_id -> integer count
|
||||
self.sent_counters = Counter()
|
||||
self.recv_counters = Counter()
|
||||
|
||||
def init(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return "FEED"
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
return COMPONENT_TYPE.CONDUIT
|
||||
|
||||
# -------------
|
||||
# Core Methods
|
||||
# -------------
|
||||
|
||||
def open(self):
|
||||
self.pull_socket = self.bind_data()
|
||||
self.feed_socket = self.bind_feed()
|
||||
|
||||
def do_work(self):
|
||||
# wait for synchronization reply from the host
|
||||
socks = dict(self.poll.poll(self.heartbeat_timeout))
|
||||
|
||||
# TODO: Abstract this out, maybe on base component
|
||||
if self.control_in in socks and socks[self.control_in] == self.zmq.POLLIN:
|
||||
msg = self.control_in.recv()
|
||||
event, payload = CONTROL_UNFRAME(msg)
|
||||
|
||||
# -- Heartbeat --
|
||||
if event == CONTROL_PROTOCOL.HEARTBEAT:
|
||||
# Heart outgoing
|
||||
heartbeat_frame = CONTROL_FRAME(
|
||||
CONTROL_PROTOCOL.OK,
|
||||
payload
|
||||
)
|
||||
self.control_out.send(heartbeat_frame)
|
||||
|
||||
# -- Soft Kill --
|
||||
elif event == CONTROL_PROTOCOL.SHUTDOWN:
|
||||
self.signal_done()
|
||||
self.shutdown()
|
||||
|
||||
# -- Hard Kill --
|
||||
elif event == CONTROL_PROTOCOL.KILL:
|
||||
self.kill()
|
||||
|
||||
|
||||
if self.pull_socket in socks and socks[self.pull_socket] == self.zmq.POLLIN:
|
||||
message = self.pull_socket.recv()
|
||||
|
||||
if message == str(CONTROL_PROTOCOL.DONE):
|
||||
self.ds_finished_counter += 1
|
||||
|
||||
if len(self.data_buffer) == self.ds_finished_counter:
|
||||
#drain any remaining messages in the buffer
|
||||
qutil.LOGGER.debug("draining feed")
|
||||
self.drain()
|
||||
self.signal_done()
|
||||
else:
|
||||
try:
|
||||
event = self.unframe(message)
|
||||
# deserialization error
|
||||
except zp.INVALID_DATASOURCE_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
try:
|
||||
self.append(event)
|
||||
self.send_next()
|
||||
|
||||
# Invalid message
|
||||
except zp.INVALID_DATASOURCE_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
def unframe(self, msg):
|
||||
return zp.DATASOURCE_UNFRAME(msg)
|
||||
|
||||
def frame(self, event):
|
||||
return zp.FEED_FRAME(event)
|
||||
|
||||
# -------------
|
||||
# Flow Control
|
||||
# -------------
|
||||
|
||||
def drain(self):
|
||||
"""
|
||||
Send all messages in the buffer.
|
||||
"""
|
||||
self.draining = True
|
||||
while self.pending_messages() > 0:
|
||||
self.send_next()
|
||||
|
||||
def send_next(self):
|
||||
"""
|
||||
Send the (chronologically) next message in the buffer.
|
||||
"""
|
||||
if not (self.is_full() or self.draining):
|
||||
return
|
||||
|
||||
event = self.next()
|
||||
if(event != None):
|
||||
self.feed_socket.send(self.frame(event), self.zmq.NOBLOCK)
|
||||
self.sent_counters[event.source_id] += 1
|
||||
self.sent_count += 1
|
||||
|
||||
def append(self, event):
|
||||
"""
|
||||
Add an event to the buffer for the source specified by
|
||||
source_id.
|
||||
"""
|
||||
self.data_buffer[event.source_id].append(event)
|
||||
self.recv_counters[event.source_id] += 1
|
||||
self.received_count += 1
|
||||
|
||||
def next(self):
|
||||
"""
|
||||
Get the next message in chronological order.
|
||||
"""
|
||||
if not(self.is_full() or self.draining):
|
||||
return
|
||||
|
||||
cur_source = None
|
||||
earliest_source = None
|
||||
earliest_event = None
|
||||
#iterate over the queues of events from all sources
|
||||
#(1 queue per datasource)
|
||||
for events in self.data_buffer.values():
|
||||
if len(events) == 0:
|
||||
continue
|
||||
cur_source = events
|
||||
first_in_list = events[0]
|
||||
if first_in_list.dt == None:
|
||||
#this is a filler event, discard
|
||||
events.pop(0)
|
||||
continue
|
||||
|
||||
if (earliest_event == None) or (first_in_list.dt <= earliest_event.dt):
|
||||
earliest_event = first_in_list
|
||||
earliest_source = cur_source
|
||||
|
||||
if earliest_event != None:
|
||||
return earliest_source.pop(0)
|
||||
|
||||
def is_full(self):
|
||||
"""
|
||||
Indicates whether the buffer has messages in buffer for
|
||||
all un-DONE, blocking sources.
|
||||
"""
|
||||
for source_id, events in self.data_buffer.iteritems():
|
||||
if len(events) == 0:
|
||||
return False
|
||||
return True
|
||||
|
||||
def pending_messages(self):
|
||||
"""
|
||||
Returns the count of all events from all sources in the
|
||||
buffer.
|
||||
"""
|
||||
total = 0
|
||||
for events in self.data_buffer.values():
|
||||
total += len(events)
|
||||
return total
|
||||
|
||||
def add_source(self, source_id):
|
||||
"""
|
||||
Add a data source to the buffer.
|
||||
"""
|
||||
self.data_buffer[source_id] = []
|
||||
|
||||
def __len__(self):
|
||||
"""
|
||||
Buffer's length is same as internal map holding separate
|
||||
sorted arrays of events keyed by source id.
|
||||
"""
|
||||
return len(self.data_buffer)
|
||||
|
||||
|
||||
class Merge(Feed):
|
||||
"""
|
||||
Merges multiple streams of events into single messages.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
Feed.__init__(self)
|
||||
|
||||
self.init()
|
||||
|
||||
def init(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return "MERGE"
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
return COMPONENT_TYPE.CONDUIT
|
||||
|
||||
def open(self):
|
||||
self.pull_socket = self.bind_merge()
|
||||
self.feed_socket = self.bind_result()
|
||||
|
||||
def next(self):
|
||||
"""Get the next merged message from the feed buffer."""
|
||||
if not (self.is_full() or self.draining):
|
||||
return
|
||||
|
||||
if self.pending_messages() == 0:
|
||||
return
|
||||
|
||||
#
|
||||
#get the raw event from the passthrough transform.
|
||||
result = self.data_buffer[zp.TRANSFORM_TYPE.PASSTHROUGH].pop(0).PASSTHROUGH
|
||||
for source, events in self.data_buffer.iteritems():
|
||||
if source == zp.TRANSFORM_TYPE.PASSTHROUGH:
|
||||
continue
|
||||
if len(events) > 0:
|
||||
cur = events.pop(0)
|
||||
result.merge(cur)
|
||||
return result
|
||||
|
||||
def unframe(self, msg):
|
||||
return zp.TRANSFORM_UNFRAME(msg)
|
||||
|
||||
def frame(self, event):
|
||||
return zp.MERGE_FRAME(event)
|
||||
|
||||
def append(self, event):
|
||||
"""
|
||||
:param event: a namedict with one entry. key is the name of the
|
||||
transform, value is the transformed value.
|
||||
Add an event to the buffer for the source specified by
|
||||
source_id.
|
||||
"""
|
||||
|
||||
self.data_buffer[event.keys()[0]].append(event)
|
||||
self.received_count += 1
|
||||
|
||||
|
||||
class BaseTransform(Component):
|
||||
"""
|
||||
Top level execution entry point for the transform
|
||||
|
||||
- connects to the feed socket to subscribe to events
|
||||
- connects to the result socket (most oftened bound by a TransformsMerge) to PUSH transforms
|
||||
- processes all messages received from feed, until DONE message received
|
||||
- pushes all transforms
|
||||
- sends DONE to result socket, closes all sockets and context
|
||||
|
||||
Parent class for feed transforms. Subclass and override transform
|
||||
method to create a new derived value from the combined feed.
|
||||
"""
|
||||
|
||||
def __init__(self, name):
|
||||
Component.__init__(self)
|
||||
|
||||
self.state = {
|
||||
'name': name
|
||||
}
|
||||
|
||||
self.init()
|
||||
|
||||
def init(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return self.state['name']
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
return COMPONENT_TYPE.CONDUIT
|
||||
|
||||
def open(self):
|
||||
"""
|
||||
Establishes zmq connections.
|
||||
"""
|
||||
#create the feed.
|
||||
self.feed_socket = self.connect_feed()
|
||||
#create the result PUSH
|
||||
self.result_socket = self.connect_merge()
|
||||
|
||||
def do_work(self):
|
||||
"""
|
||||
Loops until feed's DONE message is received:
|
||||
|
||||
- receive an event from the data feed
|
||||
- call transform (subclass' method) on event
|
||||
- send the transformed event
|
||||
|
||||
"""
|
||||
socks = dict(self.poll.poll(self.heartbeat_timeout))
|
||||
|
||||
# TODO: Abstract this out, maybe on base component
|
||||
if self.control_in in socks and socks[self.control_in] == self.zmq.POLLIN:
|
||||
msg = self.control_in.recv()
|
||||
event, payload = CONTROL_UNFRAME(msg)
|
||||
|
||||
# -- Heartbeat --
|
||||
if event == CONTROL_PROTOCOL.HEARTBEAT:
|
||||
# Heart outgoing
|
||||
heartbeat_frame = CONTROL_FRAME(
|
||||
CONTROL_PROTOCOL.OK,
|
||||
payload
|
||||
)
|
||||
self.control_out.send(heartbeat_frame)
|
||||
|
||||
# -- Soft Kill --
|
||||
elif event == CONTROL_PROTOCOL.SHUTDOWN:
|
||||
self.signal_done()
|
||||
self.shutdown()
|
||||
|
||||
# -- Hard Kill --
|
||||
elif event == CONTROL_PROTOCOL.KILL:
|
||||
self.kill()
|
||||
|
||||
if self.feed_socket in socks and socks[self.feed_socket] == self.zmq.POLLIN:
|
||||
message = self.feed_socket.recv()
|
||||
|
||||
if message == str(CONTROL_PROTOCOL.DONE):
|
||||
self.signal_done()
|
||||
return
|
||||
|
||||
try:
|
||||
event = self.unframe(message)
|
||||
except zp.INVALID_FEED_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
try:
|
||||
cur_state = self.transform(event)
|
||||
|
||||
# This is overloaded, so it can fail in all sorts of
|
||||
# unknown ways. Its best to catch it in the
|
||||
# Transformer itself.
|
||||
except Exception as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
try:
|
||||
transform_frame = self.frame(cur_state)
|
||||
except zp.INVALID_TRANSFORM_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
self.result_socket.send(transform_frame, self.zmq.NOBLOCK)
|
||||
|
||||
def frame(self, cur_state):
|
||||
return zp.TRANSFORM_FRAME(cur_state['name'], cur_state['value'])
|
||||
|
||||
def unframe(self, msg):
|
||||
return zp.FEED_UNFRAME(msg)
|
||||
|
||||
def transform(self, event):
|
||||
"""
|
||||
Must return the transformed value as a map with::
|
||||
|
||||
{name:"name of new transform", value: "value of new field"}
|
||||
|
||||
Transforms run in parallel and results are merged into a single map, so
|
||||
transform names must be unique. Best practice is to use the self.state
|
||||
object initialized from the transform configuration, and only set the
|
||||
transformed value::
|
||||
|
||||
self.state['value'] = transformed_value
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class PassthroughTransform(BaseTransform):
|
||||
"""
|
||||
A bypass transform which is also an identity transform::
|
||||
|
||||
+-------+
|
||||
+---| f |--->
|
||||
+-------+
|
||||
+------id------->
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
BaseTransform.__init__(self, "PASSTHROUGH")
|
||||
self.init()
|
||||
|
||||
def init(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
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 {'name':zp.TRANSFORM_TYPE.PASSTHROUGH, 'value': zp.FEED_FRAME(event) }
|
||||
|
||||
|
||||
class DataSource(Component):
|
||||
"""
|
||||
Baseclass for data sources. Subclass and implement send_all - usually this
|
||||
means looping through all records in a store, converting to a dict, and
|
||||
calling send(map).
|
||||
|
||||
Every datasource has a dict property to hold filters::
|
||||
- key -- name of the filter, e.g. SID
|
||||
- value -- a primitive representing the filter. e.g. a list of ints.
|
||||
|
||||
Modify the datasource's filters via the set_filter(name, value)
|
||||
"""
|
||||
def __init__(self, source_id):
|
||||
Component.__init__(self)
|
||||
|
||||
self.id = source_id
|
||||
self.init()
|
||||
self.filter = {}
|
||||
|
||||
def init(self):
|
||||
self.cur_event = None
|
||||
|
||||
def set_filter(self, name, value):
|
||||
self.filter[name] = value
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return self.id
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
return COMPONENT_TYPE.SOURCE
|
||||
|
||||
def open(self):
|
||||
self.data_socket = self.connect_data()
|
||||
|
||||
def send(self, event):
|
||||
"""
|
||||
Emit data.
|
||||
"""
|
||||
assert isinstance(event, zp.namedict)
|
||||
|
||||
event['source_id'] = self.get_id
|
||||
event['type'] = self.get_type
|
||||
|
||||
try:
|
||||
ds_frame = self.frame(event)
|
||||
except zp.INVALID_DATASOURCE_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
self.data_socket.send(ds_frame)
|
||||
|
||||
def frame(self, event):
|
||||
return zp.DATASOURCE_FRAME(event)
|
||||
@@ -1,18 +1,18 @@
|
||||
import zmq
|
||||
import time
|
||||
import gevent
|
||||
import itertools
|
||||
# pyzmq
|
||||
import zmq
|
||||
import logging
|
||||
import gevent_zeromq
|
||||
|
||||
from collections import OrderedDict
|
||||
|
||||
from protocol import CONTROL_PROTOCOL, CONTROL_FRAME, \
|
||||
from zipline.protocol import CONTROL_PROTOCOL, CONTROL_FRAME, \
|
||||
CONTROL_UNFRAME, CONTROL_STATES, INVALID_CONTROL_FRAME \
|
||||
|
||||
states = CONTROL_STATES
|
||||
|
||||
from gpoll import _Poller as GeventPoller
|
||||
from zipline.utils.gpoll import _Poller as GeventPoller
|
||||
|
||||
# Roll Call ( Discovery )
|
||||
# -----------------------
|
||||
@@ -159,7 +159,7 @@ class Controller(object):
|
||||
debug = False
|
||||
period = 1
|
||||
|
||||
def __init__(self, pub_socket, route_socket, logging = None):
|
||||
def __init__(self, pub_socket, route_socket, logger = None):
|
||||
|
||||
self.context = None
|
||||
self.zmq = None
|
||||
@@ -182,11 +182,11 @@ class Controller(object):
|
||||
|
||||
self.error_replay = OrderedDict()
|
||||
|
||||
if logging:
|
||||
self.logging = logging
|
||||
if logger:
|
||||
self.logging = logger
|
||||
else:
|
||||
import util as qutil
|
||||
self.logging = qutil.LOGGER
|
||||
default_logger = logging.getLogger('ZiplineLogger')
|
||||
self.logging = default_logger
|
||||
|
||||
def init_zmq(self, flavor):
|
||||
|
||||
|
||||
@@ -40,8 +40,8 @@ Performance Tracking
|
||||
| | through all the events delivered to this tracker. |
|
||||
| | For details look at the comments for |
|
||||
| | :py:meth:`zipline.finance.risk.RiskMetrics.to_dict`|
|
||||
+-----------------+----------------------------------------------------+
|
||||
| exceeded_max_ | True if the simulation was stopped because single |
|
||||
+-----------------+----------------------------------------------------+
|
||||
| exceeded_max_ | True if the simulation was stopped because single |
|
||||
| loss | day losses exceeded the max_drawdown stipulated in |
|
||||
| | trading_environment. |
|
||||
+-----------------+----------------------------------------------------+
|
||||
@@ -110,6 +110,8 @@ Performance Period
|
||||
|
||||
|
||||
"""
|
||||
|
||||
import logging
|
||||
import datetime
|
||||
import pytz
|
||||
import msgpack
|
||||
@@ -118,10 +120,11 @@ import math
|
||||
|
||||
import zmq
|
||||
|
||||
import zipline.util as qutil
|
||||
import zipline.protocol as zp
|
||||
import zipline.finance.risk as risk
|
||||
|
||||
LOGGER = logging.getLogger('ZiplineLogger')
|
||||
|
||||
class PerformanceTracker():
|
||||
"""
|
||||
Tracks the performance of the zipline as it is running in
|
||||
@@ -188,7 +191,7 @@ class PerformanceTracker():
|
||||
)
|
||||
|
||||
def get_portfolio(self):
|
||||
return self.cumulative_performance.to_namedict()
|
||||
return self.cumulative_performance.to_ndict()
|
||||
|
||||
def publish_to(self, zmq_socket, context=None):
|
||||
"""
|
||||
@@ -228,7 +231,7 @@ class PerformanceTracker():
|
||||
if self.exceeded_max_loss:
|
||||
return
|
||||
|
||||
assert isinstance(event, zp.namedict)
|
||||
assert isinstance(event, zp.ndict)
|
||||
self.event_count += 1
|
||||
|
||||
if(event.dt >= self.market_close):
|
||||
@@ -280,8 +283,8 @@ class PerformanceTracker():
|
||||
returns = self.todays_performance.returns
|
||||
max_dd = -1 * self.trading_environment.max_drawdown
|
||||
if returns < max_dd:
|
||||
qutil.LOGGER.info(str(returns) + " broke through " + str(max_dd))
|
||||
qutil.LOGGER.info("Exceeded max drawdown.")
|
||||
LOGGER.info(str(returns) + " broke through " + str(max_dd))
|
||||
LOGGER.info("Exceeded max drawdown.")
|
||||
# mark the perf period with max loss flag,
|
||||
# so it shows up in the update, but don't end the test
|
||||
# here. Let the update go out before stopping
|
||||
@@ -316,8 +319,8 @@ class PerformanceTracker():
|
||||
"""
|
||||
|
||||
log_msg = "Simulated {n} trading days out of {m}."
|
||||
qutil.LOGGER.info(log_msg.format(n=self.day_count, m=self.total_days))
|
||||
qutil.LOGGER.info("first open: {d}".format(d=self.trading_environment.first_open))
|
||||
LOGGER.info(log_msg.format(n=self.day_count, m=self.total_days))
|
||||
LOGGER.info("first open: {d}".format(d=self.trading_environment.first_open))
|
||||
|
||||
# the stream will end on the last trading day, but will not trigger
|
||||
# an end of day, so we trigger the final market close here.
|
||||
@@ -332,7 +335,7 @@ class PerformanceTracker():
|
||||
)
|
||||
|
||||
if self.result_stream:
|
||||
qutil.LOGGER.info("about to stream the risk report...")
|
||||
LOGGER.info("about to stream the risk report...")
|
||||
risk_dict = self.risk_report.to_dict()
|
||||
|
||||
msg = zp.RISK_FRAME(risk_dict)
|
||||
@@ -518,18 +521,18 @@ class PerformancePeriod():
|
||||
|
||||
return rval
|
||||
|
||||
def to_namedict(self):
|
||||
def to_ndict(self):
|
||||
"""
|
||||
Creates a namedict representing the state of this perfomance period.
|
||||
Creates a ndict representing the state of this perfomance period.
|
||||
Properties are the same as the results of to_dict. See header comments
|
||||
for a detailed description.
|
||||
|
||||
"""
|
||||
positions = self.get_positions(namedicted=True)
|
||||
positions = self.get_positions(ndicted=True)
|
||||
|
||||
positions = zp.namedict(positions)
|
||||
positions = zp.ndict(positions)
|
||||
|
||||
return zp.namedict({
|
||||
return zp.ndict({
|
||||
'ending_value' : self.ending_value,
|
||||
'capital_used' : self.period_capital_used,
|
||||
'starting_value' : self.starting_value,
|
||||
@@ -542,12 +545,12 @@ class PerformancePeriod():
|
||||
'transactions' : self.processed_transactions
|
||||
})
|
||||
|
||||
def get_positions(self, namedicted=False):
|
||||
def get_positions(self, ndicted=False):
|
||||
positions = {}
|
||||
for sid, pos in self.positions.iteritems():
|
||||
cur = pos.to_dict()
|
||||
if namedicted:
|
||||
positions[sid] = zp.namedict(cur)
|
||||
if ndicted:
|
||||
positions[sid] = zp.ndict(cur)
|
||||
else:
|
||||
positions[sid] = cur
|
||||
|
||||
|
||||
@@ -36,14 +36,15 @@ Risk Report
|
||||
|
||||
"""
|
||||
|
||||
import logging
|
||||
import datetime
|
||||
import math
|
||||
import pytz
|
||||
import numpy as np
|
||||
import numpy.linalg as la
|
||||
import zipline.util as qutil
|
||||
import zipline.protocol as zp
|
||||
|
||||
LOGGER = logging.getLogger('ZiplineLogger')
|
||||
|
||||
def advance_by_months(dt, jump_in_months):
|
||||
month = dt.month + jump_in_months
|
||||
@@ -243,7 +244,7 @@ class RiskMetrics():
|
||||
cur_return = math.log(1.0 + r) + cur_return
|
||||
#this is a guard for a single day returning -100%
|
||||
except ValueError:
|
||||
qutil.LOGGER.warn("{cur} return, zeroing the returns".format(cur=cur_return))
|
||||
LOGGER.warn("{cur} return, zeroing the returns".format(cur=cur_return))
|
||||
cur_return = 0.0
|
||||
compounded_returns.append(cur_return)
|
||||
|
||||
|
||||
+14
-17
@@ -5,11 +5,12 @@ import datetime
|
||||
import random
|
||||
import pytz
|
||||
|
||||
import zipline.messaging as zm
|
||||
from zipline.components import DataSource
|
||||
from zipline.utils import ndict
|
||||
|
||||
import zipline.protocol as zp
|
||||
|
||||
|
||||
class TradeDataSource(zm.DataSource):
|
||||
class TradeDataSource(DataSource):
|
||||
|
||||
def send(self, event):
|
||||
"""
|
||||
@@ -18,19 +19,19 @@ class TradeDataSource(zm.DataSource):
|
||||
:py:func: `zipline.protocol.TRADE_FRAME`
|
||||
:rtype: None
|
||||
"""
|
||||
|
||||
|
||||
event.source_id = self.get_id
|
||||
if event.sid in self.filter['SID']:
|
||||
if event.sid in self.filter['SID']:
|
||||
message = zp.DATASOURCE_FRAME(event)
|
||||
else:
|
||||
blank = zp.namedict({
|
||||
blank = ndict({
|
||||
"type" : zp.DATASOURCE_TYPE.TRADE,
|
||||
"source_id" : self.get_id
|
||||
})
|
||||
message = zp.DATASOURCE_FRAME(blank)
|
||||
|
||||
|
||||
self.data_socket.send(message)
|
||||
|
||||
|
||||
|
||||
class RandomEquityTrades(TradeDataSource):
|
||||
"""
|
||||
@@ -38,7 +39,7 @@ class RandomEquityTrades(TradeDataSource):
|
||||
"""
|
||||
|
||||
def __init__(self, sid, source_id, count):
|
||||
zm.DataSource.__init__(self, source_id)
|
||||
DataSource.__init__(self, source_id)
|
||||
self.count = count
|
||||
self.incr = 0
|
||||
self.sid = sid
|
||||
@@ -58,7 +59,7 @@ class RandomEquityTrades(TradeDataSource):
|
||||
self.price = self.price + random.uniform(-0.05, 0.05)
|
||||
volume = random.randrange(100,10000,100)
|
||||
|
||||
event = zp.namedict({
|
||||
event = zp.ndict({
|
||||
"type" : zp.DATASOURCE_TYPE.TRADE,
|
||||
"sid" : self.sid,
|
||||
"price" : self.price,
|
||||
@@ -67,7 +68,6 @@ class RandomEquityTrades(TradeDataSource):
|
||||
})
|
||||
self.send(event)
|
||||
self.incr += 1
|
||||
|
||||
|
||||
|
||||
class SpecificEquityTrades(TradeDataSource):
|
||||
@@ -77,7 +77,7 @@ class SpecificEquityTrades(TradeDataSource):
|
||||
|
||||
def __init__(self, source_id, event_list):
|
||||
"""
|
||||
:param event_list: should be a chronologically ordered list of
|
||||
:param event_list: should be a chronologically ordered list of
|
||||
dictionaries in the following form:
|
||||
|
||||
event = {
|
||||
@@ -87,21 +87,18 @@ class SpecificEquityTrades(TradeDataSource):
|
||||
'volume' : integer for volume
|
||||
}
|
||||
"""
|
||||
zm.DataSource.__init__(self, source_id)
|
||||
DataSource.__init__(self, source_id)
|
||||
self.event_list = event_list
|
||||
self.count = 0
|
||||
|
||||
def get_type(self):
|
||||
zp.COMPONENT_TYPE.SOURCE
|
||||
|
||||
|
||||
def do_work(self):
|
||||
if(len(self.event_list) == 0):
|
||||
self.signal_done()
|
||||
return
|
||||
|
||||
event = self.event_list.pop(0)
|
||||
self.send(zp.namedict(event))
|
||||
self.send(zp.ndict(event))
|
||||
self.count +=1
|
||||
|
||||
|
||||
|
||||
+12
-11
@@ -1,3 +1,4 @@
|
||||
import logging
|
||||
import datetime
|
||||
import pytz
|
||||
import math
|
||||
@@ -6,14 +7,12 @@ import time
|
||||
from collections import Counter
|
||||
|
||||
# from gevent.select import select
|
||||
from zmq.core.poll import select
|
||||
|
||||
import zipline.messaging as qmsg
|
||||
import zipline.util as qutil
|
||||
from zipline.core import Component
|
||||
import zipline.protocol as zp
|
||||
import zipline.finance.performance as perf
|
||||
|
||||
from zipline.protocol_utils import Enum, ndict
|
||||
from zipline.utils.protocol_utils import Enum, ndict
|
||||
|
||||
# the simulation style enumerates the available transaction simulation
|
||||
# strategies.
|
||||
@@ -24,10 +23,12 @@ SIMULATION_STYLE = Enum(
|
||||
'NOOP'
|
||||
)
|
||||
|
||||
class TradeSimulationClient(qmsg.Component):
|
||||
LOGGER = logging.getLogger('ZiplineLogger')
|
||||
|
||||
class TradeSimulationClient(Component):
|
||||
|
||||
def __init__(self, trading_environment, sim_style):
|
||||
qmsg.Component.__init__(self)
|
||||
Component.__init__(self)
|
||||
self.received_count = 0
|
||||
self.prev_dt = None
|
||||
self.event_queue = None
|
||||
@@ -89,7 +90,7 @@ class TradeSimulationClient(qmsg.Component):
|
||||
self.finish_simulation()
|
||||
|
||||
def finish_simulation(self):
|
||||
qutil.LOGGER.info("Client is DONE!")
|
||||
LOGGER.info("Client is DONE!")
|
||||
# signal the performance tracker that the simulation has
|
||||
# ended. Perf will internally calculate the full risk report.
|
||||
self.perf.handle_simulation_end()
|
||||
@@ -158,7 +159,7 @@ class TradeSimulationClient(qmsg.Component):
|
||||
return self.connect_push_socket(self.addresses['order_address'])
|
||||
|
||||
def order(self, sid, amount):
|
||||
order = zp.namedict({
|
||||
order = zp.ndict({
|
||||
'dt':self.current_dt,
|
||||
'sid':sid,
|
||||
'amount':amount
|
||||
@@ -213,7 +214,7 @@ class TransactionSimulator(object):
|
||||
log = "requested to trade zero shares of {sid}".format(
|
||||
sid=event.sid
|
||||
)
|
||||
qutil.LOGGER.debug(log)
|
||||
LOGGER.debug(log)
|
||||
return
|
||||
|
||||
if(not self.open_orders.has_key(event.sid)):
|
||||
@@ -337,7 +338,7 @@ for orders:
|
||||
event=str(event),
|
||||
orders=str(orders)
|
||||
)
|
||||
qutil.LOGGER.warn(warning)
|
||||
LOGGER.warn(warning)
|
||||
return None
|
||||
|
||||
|
||||
@@ -350,7 +351,7 @@ for orders:
|
||||
'commission' : self.commission * amount * direction,
|
||||
'source_id' : zp.FINANCE_COMPONENT.TRANSACTION_SIM
|
||||
}
|
||||
return zp.namedict(txn)
|
||||
return zp.ndict(txn)
|
||||
|
||||
|
||||
class TradingEnvironment(object):
|
||||
|
||||
+70
-78
@@ -1,20 +1,20 @@
|
||||
"""
|
||||
Ziplines are composed of multiple components connected by asynchronous
|
||||
messaging. All ziplines follow a general topology of parallel sources,
|
||||
datetimestamp serialization, parallel transformations, and finally sinks.
|
||||
Furthermore, many ziplines have common needs. For example, all trade
|
||||
simulations require a
|
||||
Ziplines are composed of multiple components connected by asynchronous
|
||||
messaging. All ziplines follow a general topology of parallel sources,
|
||||
datetimestamp serialization, parallel transformations, and finally sinks.
|
||||
Furthermore, many ziplines have common needs. For example, all trade
|
||||
simulations require a
|
||||
:py:class:`~zipline.finance.trading.TradeSimulationClient`.
|
||||
|
||||
To establish best practices and minimize code replication, the lines module
|
||||
To establish best practices and minimize code replication, the lines module
|
||||
provides complete zipline topologies. You can extend any zipline without
|
||||
the need to extend the class. Simply instantiate any additional components
|
||||
that you would like included in the zipline, and add them to the zipline
|
||||
before invoking simulate.
|
||||
that you would like included in the zipline, and add them to the zipline
|
||||
before invoking simulate.
|
||||
|
||||
|
||||
|
||||
Here is a diagram of the SimulatedTrading zipline:
|
||||
|
||||
|
||||
|
||||
+----------------------+ +------------------------+
|
||||
| Trade History | | (DataSource added |
|
||||
@@ -60,62 +60,54 @@ before invoking simulate.
|
||||
+---------------------------------+
|
||||
"""
|
||||
|
||||
import mock
|
||||
import pytz
|
||||
import logging
|
||||
|
||||
from datetime import datetime, timedelta
|
||||
from collections import defaultdict
|
||||
import zipline.utils.factory as factory
|
||||
|
||||
from nose.tools import timed
|
||||
from zipline.components import DataSource
|
||||
from zipline.transforms import BaseTransform
|
||||
|
||||
import zipline.test.factory as factory
|
||||
import zipline.util as qutil
|
||||
import zipline.finance.risk as risk
|
||||
import zipline.protocol as zp
|
||||
import zipline.finance.performance as perf
|
||||
import zipline.messaging as zmsg
|
||||
|
||||
from zipline.test.algorithms import TestAlgorithm
|
||||
from zipline.sources import SpecificEquityTrades
|
||||
from zipline.test_algorithms import TestAlgorithm
|
||||
from zipline.finance.trading import TradeSimulationClient
|
||||
from zipline.simulator import AddressAllocator, Simulator
|
||||
from zipline.monitor import Controller
|
||||
from zipline.core.devsimulator import Simulator
|
||||
from zipline.core.monitor import Controller
|
||||
from zipline.finance.trading import SIMULATION_STYLE
|
||||
|
||||
LOGGER = logging.getLogger('ZiplineLogger')
|
||||
|
||||
class SimulatedTrading(object):
|
||||
"""
|
||||
Zipline with::
|
||||
|
||||
|
||||
- _no_ data sources.
|
||||
- Trade simulation client, which is available to send callbacks on
|
||||
events and also accept orders to be simulated.
|
||||
- An order data source, which will receive orders from the trade
|
||||
simulation client, and feed them into the event stream to be
|
||||
simulation client, and feed them into the event stream to be
|
||||
serialized and order alongside all other data source events.
|
||||
- transaction simulation transformation, which receives the order
|
||||
events and estimates a theoretical execution price and volume.
|
||||
|
||||
|
||||
All components in this zipline are subject to heartbeat checks and
|
||||
a control monitor, which can kill the entire zipline in the event of
|
||||
exceptions in one of the components or an external request to end the
|
||||
simulation.
|
||||
"""
|
||||
|
||||
|
||||
def __init__(self, **config):
|
||||
"""
|
||||
:param config: a dict with the following required properties::
|
||||
|
||||
|
||||
- algorithm: a class that follows the algorithm protocol. See
|
||||
:py:meth:`zipline.finance.trading.TradingSimulationClient.add_algorithm`
|
||||
:py:meth:`zipline.finance.trading.TradingSimulationClient.add_algorithm
|
||||
for details.
|
||||
- trading_environment: an instance of
|
||||
:py:class:`zipline.trading.TradingEnvironment`
|
||||
- allocator: an instance of
|
||||
- allocator: an instance of
|
||||
:py:class:`zipline.simulator.AddressAllocator`
|
||||
- simulator_class: a :py:class:`zipline.messaging.ComponentHost`
|
||||
- simulator_class: a :py:class:`zipline.core.host.ComponentHost`
|
||||
subclass (not an instance)
|
||||
- simulation_style: optional parameter that configures the
|
||||
- simulation_style: optional parameter that configures the
|
||||
:py:class:`zipline.finance.trading.TransactionSimulator`. Expects
|
||||
a SIMULATION_STYLE as defined in :py:mod:`zipline.finance.trading`
|
||||
"""
|
||||
@@ -124,10 +116,10 @@ class SimulatedTrading(object):
|
||||
self.allocator = config['allocator']
|
||||
self.trading_environment = config['trading_environment']
|
||||
self.sim_style = config.get('simulation_style')
|
||||
|
||||
|
||||
self.leased_sockets = []
|
||||
self.sim_context = None
|
||||
|
||||
|
||||
sockets = self.allocate_sockets(8)
|
||||
addresses = {
|
||||
'sync_address' : sockets[0],
|
||||
@@ -141,75 +133,75 @@ class SimulatedTrading(object):
|
||||
self.con = Controller(
|
||||
sockets[6],
|
||||
sockets[7],
|
||||
logging = qutil.LOGGER
|
||||
logger = LOGGER
|
||||
)
|
||||
|
||||
|
||||
self.con.cancel_socket = self.allocator.lease(1)[0]
|
||||
|
||||
# TODO: Not freeform
|
||||
self.con.manage(
|
||||
'freeform'
|
||||
)
|
||||
|
||||
|
||||
self.started = False
|
||||
|
||||
|
||||
self.sim = config['simulator_class'](addresses)
|
||||
|
||||
|
||||
self.clients = {}
|
||||
self.trading_client = TradeSimulationClient(
|
||||
self.trading_environment,
|
||||
self.sim_style
|
||||
)
|
||||
self.add_client(self.trading_client)
|
||||
|
||||
|
||||
# setup all sources
|
||||
self.sources = {}
|
||||
#self.order_source = OrderDataSource()
|
||||
#self.add_source(self.order_source)
|
||||
|
||||
|
||||
#setup transforms
|
||||
#self.transaction_sim = TransactionSimulator(self.sim_style)
|
||||
self.transforms = {}
|
||||
#self.add_transform(self.transaction_sim)
|
||||
|
||||
|
||||
self.sim.register_controller( self.con )
|
||||
self.sim.on_done = self.shutdown()
|
||||
|
||||
|
||||
|
||||
|
||||
self.trading_client.set_algorithm(self.algorithm)
|
||||
|
||||
|
||||
|
||||
@staticmethod
|
||||
def create_test_zipline(**config):
|
||||
"""
|
||||
:param config: A configuration object that is a dict with:
|
||||
|
||||
|
||||
- environment - a \
|
||||
:py:class:`zipline.finance.trading.TradingEnvironment`
|
||||
- allocator - a :py:class:`zipline.simulator.AddressAllocator`
|
||||
- sid - an integer, which will be used as the security ID.
|
||||
- sid - an integer, which will be used as the security ID.
|
||||
- order_count - the number of orders the test algo will place,
|
||||
defaults to 100
|
||||
- order_amount - the number of shares per order, defaults to 100
|
||||
- trade_count - the number of trades to simulate, defaults to 101
|
||||
to ensure all orders are processed.
|
||||
- simulator_class - optional parameter that provides an alternative
|
||||
- simulator_class - optional parameter that provides an alternative
|
||||
subclass of ComponentHost to hold the whole zipline. Defaults to
|
||||
:py:class:`zipline.simulator.Simulator`
|
||||
:py:class:`zipline.simulator.Simulator`
|
||||
- algorithm - optional parameter providing an algorithm. defaults
|
||||
to :py:class:`zipline.test.algorithms.TestAlgorithm`
|
||||
- trade_source - optional parameter to specify trades, if present.
|
||||
If not present :py:class:`ziplien.sources.SpecificEquityTrades`
|
||||
If not present :py:class:`ziplien.sources.SpecificEquityTrades`
|
||||
is the source, with daily frequency in trades.
|
||||
- simulation_style: optional parameter that configures the
|
||||
- simulation_style: optional parameter that configures the
|
||||
:py:class:`zipline.finance.trading.TransactionSimulator`. Expects
|
||||
a SIMULATION_STYLE as defined in :py:mod:`zipline.finance.trading`
|
||||
"""
|
||||
assert isinstance(config, dict)
|
||||
|
||||
|
||||
allocator = config['allocator']
|
||||
sid = config['sid']
|
||||
|
||||
|
||||
#--------------------
|
||||
# Trading Environment
|
||||
#--------------------
|
||||
@@ -217,33 +209,33 @@ class SimulatedTrading(object):
|
||||
trading_environment = config['environment']
|
||||
else:
|
||||
trading_environment = factory.create_trading_environment()
|
||||
|
||||
|
||||
if config.has_key('order_count'):
|
||||
order_count = config['order_count']
|
||||
else:
|
||||
order_count = 100
|
||||
|
||||
|
||||
if config.has_key('order_amount'):
|
||||
order_amount = config['order_amount']
|
||||
else:
|
||||
order_amount = 100
|
||||
|
||||
|
||||
if config.has_key('trade_count'):
|
||||
trade_count = config['trade_count']
|
||||
else:
|
||||
# to ensure all orders are filled, we provide one more
|
||||
# trade than order
|
||||
trade_count = 101
|
||||
|
||||
|
||||
if config.has_key('simulator_class'):
|
||||
simulator_class = config['simulator_class']
|
||||
else:
|
||||
simulator_class = Simulator
|
||||
|
||||
|
||||
simulation_style = config.get('simulation_style')
|
||||
if not simulation_style:
|
||||
simulation_style = SIMULATION_STYLE.FIXED_SLIPPAGE
|
||||
|
||||
|
||||
#-------------------
|
||||
# Trade Source
|
||||
#-------------------
|
||||
@@ -283,41 +275,41 @@ class SimulatedTrading(object):
|
||||
zipline.add_source(trade_source)
|
||||
|
||||
return zipline
|
||||
|
||||
|
||||
def add_source(self, source):
|
||||
"""
|
||||
Adds the source to the zipline, sets the sid filter of the
|
||||
source to the algorithm's sid filter.
|
||||
"""
|
||||
assert isinstance(source, zmsg.DataSource)
|
||||
self.check_started()
|
||||
assert isinstance(source, DataSource)
|
||||
self.check_started()
|
||||
source.set_filter('SID', self.algorithm.get_sid_filter())
|
||||
self.sim.register_components([source])
|
||||
self.sources[source.get_id] = source
|
||||
|
||||
|
||||
def add_transform(self, transform):
|
||||
assert isinstance(transform, zmsg.BaseTransform)
|
||||
assert isinstance(transform, BaseTransform)
|
||||
self.check_started()
|
||||
self.sim.register_components([transform])
|
||||
self.transforms[transform.get_id] = transform
|
||||
|
||||
|
||||
def add_client(self, client):
|
||||
assert isinstance(client, TradeSimulationClient)
|
||||
self.check_started()
|
||||
self.sim.register_components([client])
|
||||
self.clients[client.get_id] = client
|
||||
|
||||
|
||||
def check_started(self):
|
||||
if self.started:
|
||||
raise ZiplineException("TradeSimulation", "You cannot add \
|
||||
components after the simulation has begun.")
|
||||
|
||||
|
||||
def get_cumulative_performance(self):
|
||||
return self.trading_client.perf.cumulative_performance.to_dict()
|
||||
|
||||
|
||||
def publish_to(self, result_socket):
|
||||
self.trading_client.perf.publish_to(result_socket)
|
||||
|
||||
|
||||
def allocate_sockets(self, n):
|
||||
"""
|
||||
Allocate sockets local to this line, track them so
|
||||
@@ -331,7 +323,7 @@ class SimulatedTrading(object):
|
||||
|
||||
self.leased_sockets.extend(leased)
|
||||
return leased
|
||||
|
||||
|
||||
def simulate(self, blocking=False):
|
||||
self.started = True
|
||||
self.sim_context = self.sim.simulate()
|
||||
@@ -341,11 +333,11 @@ class SimulatedTrading(object):
|
||||
def shutdown(self):
|
||||
pass
|
||||
#self.allocator.reaquire(*self.leased_sockets)
|
||||
|
||||
|
||||
#--------------------------------
|
||||
# Component property accessors
|
||||
#--------------------------------
|
||||
|
||||
|
||||
def get_positions(self):
|
||||
"""
|
||||
returns current positions as a dict. draws from the cumulative
|
||||
@@ -354,14 +346,14 @@ class SimulatedTrading(object):
|
||||
perf = self.trading_client.perf.cumulative_performance
|
||||
positions = perf.get_positions()
|
||||
return positions
|
||||
|
||||
|
||||
class ZiplineException(Exception):
|
||||
def __init__(self, zipline_name, msg):
|
||||
self.name = zipline_name
|
||||
self.message = msg
|
||||
|
||||
|
||||
def __str__(self):
|
||||
return "Unexpected exception {line}: {msg}".format(
|
||||
line=self.name,
|
||||
line=self.name,
|
||||
msg=self.message
|
||||
)
|
||||
|
||||
@@ -1,636 +0,0 @@
|
||||
"""
|
||||
Commonly used messaging components.
|
||||
"""
|
||||
|
||||
import datetime
|
||||
|
||||
from collections import Counter
|
||||
|
||||
import zipline.util as qutil
|
||||
from zipline.component import Component
|
||||
import zipline.protocol as zp
|
||||
from zipline.protocol import CONTROL_PROTOCOL, COMPONENT_TYPE, \
|
||||
COMPONENT_STATE, CONTROL_FRAME, CONTROL_UNFRAME
|
||||
|
||||
class ComponentHost(Component):
|
||||
"""
|
||||
Components that can launch multiple sub-components, synchronize their
|
||||
start, and then wait for all components to be finished.
|
||||
"""
|
||||
|
||||
def __init__(self, addresses):
|
||||
Component.__init__(self)
|
||||
self.addresses = addresses
|
||||
self.running = False
|
||||
|
||||
self.init()
|
||||
|
||||
def init(self):
|
||||
assert hasattr(self, 'zmq_flavor'), \
|
||||
""" You must specify a flavor of ZeroMQ for all
|
||||
ComponentHost subclasses. """
|
||||
|
||||
# Component Registry, keyed by get_id
|
||||
# ----------------------
|
||||
self.components = {}
|
||||
# ----------------------
|
||||
# Internal Registry, keyed by guid
|
||||
self._components = {}
|
||||
# ----------------------
|
||||
|
||||
self.sync_register = {}
|
||||
self.timeout = datetime.timedelta(seconds=60)
|
||||
|
||||
self.feed = Feed()
|
||||
self.merge = Merge()
|
||||
self.passthrough = PassthroughTransform()
|
||||
self.controller = None
|
||||
|
||||
#register the feed and the merge
|
||||
self.register_components([self.feed, self.merge, self.passthrough])
|
||||
|
||||
def register_controller(self, controller):
|
||||
"""
|
||||
Add the given components to the registry. Establish
|
||||
communication with them.
|
||||
"""
|
||||
if self.controller != None:
|
||||
raise Exception("There can be only one!")
|
||||
|
||||
self.controller = controller
|
||||
self.controller.zmq_flavor = self.zmq_flavor
|
||||
|
||||
# Propogate the controller to all the subcomponents
|
||||
for component in self.components.itervalues():
|
||||
component.controller = controller
|
||||
|
||||
def register_components(self, components):
|
||||
"""
|
||||
Add the given components to the registry. Establish
|
||||
communication with them.
|
||||
"""
|
||||
assert isinstance(components, list)
|
||||
for component in components:
|
||||
|
||||
component.addresses = self.addresses
|
||||
component.controller = self.controller
|
||||
|
||||
# Hosts share their zmq flavor with hosted components
|
||||
component.zmq_flavor = self.zmq_flavor
|
||||
|
||||
self._components[component.guid] = component
|
||||
self.components[component.get_id] = component
|
||||
self.sync_register[component.get_id] = datetime.datetime.utcnow()
|
||||
|
||||
if isinstance(component, DataSource):
|
||||
self.feed.add_source(component.get_id)
|
||||
if isinstance(component, BaseTransform):
|
||||
self.merge.add_source(component.get_id)
|
||||
|
||||
def unregister_component(self, component_id):
|
||||
del self.components[component_id]
|
||||
del self.sync_register[component_id]
|
||||
|
||||
def setup_sync(self):
|
||||
"""
|
||||
Setup the sync socket and poller. ( Bind )
|
||||
"""
|
||||
qutil.LOGGER.debug("Connecting sync server.")
|
||||
|
||||
self.sync_socket = self.context.socket(self.zmq.REP)
|
||||
self.sync_socket.bind(self.addresses['sync_address'])
|
||||
|
||||
self.sync_poller = self.zmq_poller()
|
||||
self.sync_poller.register(self.sync_socket, self.zmq.POLLIN)
|
||||
|
||||
self.sockets.append(self.sync_socket)
|
||||
|
||||
def open(self):
|
||||
for component in self.components.values():
|
||||
self.launch_component(component)
|
||||
self.launch_controller()
|
||||
|
||||
def is_running(self):
|
||||
"""
|
||||
DEPRECATED, left in for compatability for now.
|
||||
"""
|
||||
|
||||
cur_time = datetime.datetime.utcnow()
|
||||
|
||||
if len(self.components) == 0:
|
||||
qutil.LOGGER.info("Component register is empty.")
|
||||
return False
|
||||
|
||||
return True
|
||||
|
||||
def loop(self, lockstep=True):
|
||||
|
||||
while self.is_running():
|
||||
# wait for synchronization request at start, and DONE at end.
|
||||
# don't timeout.
|
||||
socks = dict(self.sync_poller.poll())
|
||||
|
||||
if self.sync_socket in socks and socks[self.sync_socket] == self.zmq.POLLIN:
|
||||
msg = self.sync_socket.recv()
|
||||
|
||||
try:
|
||||
parts = msg.split(':')
|
||||
sync_id, status = parts
|
||||
except ValueError as exc:
|
||||
self.signal_exception(exc)
|
||||
|
||||
if status == str(CONTROL_PROTOCOL.DONE): # TODO: other way around
|
||||
#qutil.LOGGER.debug("{id} is DONE".format(id=sync_id))
|
||||
self.unregister_component(sync_id)
|
||||
self.state_flag = COMPONENT_STATE.DONE
|
||||
else:
|
||||
self.sync_register[sync_id] = datetime.datetime.utcnow()
|
||||
|
||||
#qutil.LOGGER.info("confirmed {id}".format(id=msg))
|
||||
# send synchronization reply
|
||||
self.sync_socket.send('ack', self.zmq.NOBLOCK)
|
||||
|
||||
# ------------------
|
||||
# Simulation Control
|
||||
# ------------------
|
||||
|
||||
def launch_controller(self, controller):
|
||||
raise NotImplementedError
|
||||
|
||||
def launch_component(self, component):
|
||||
raise NotImplementedError
|
||||
|
||||
def teardown_component(self, component):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class Feed(Component):
|
||||
"""
|
||||
Connects to N PULL sockets, publishing all messages received to a PUB
|
||||
socket. Published messages are guaranteed to be in chronological order
|
||||
based on message property dt. Expects to be instantiated in one execution
|
||||
context (thread, process, etc) and run in another.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
Component.__init__(self)
|
||||
|
||||
self.sent_count = 0
|
||||
self.received_count = 0
|
||||
self.draining = False
|
||||
self.ds_finished_counter = 0
|
||||
|
||||
# Depending on the size of this, might want to use a data
|
||||
# structure with better asymptotics.
|
||||
self.data_buffer = {}
|
||||
|
||||
# source_id -> integer count
|
||||
self.sent_counters = Counter()
|
||||
self.recv_counters = Counter()
|
||||
|
||||
def init(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return "FEED"
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
return COMPONENT_TYPE.CONDUIT
|
||||
|
||||
# -------------
|
||||
# Core Methods
|
||||
# -------------
|
||||
|
||||
def open(self):
|
||||
self.pull_socket = self.bind_data()
|
||||
self.feed_socket = self.bind_feed()
|
||||
|
||||
def do_work(self):
|
||||
# wait for synchronization reply from the host
|
||||
socks = dict(self.poll.poll(self.heartbeat_timeout))
|
||||
|
||||
# TODO: Abstract this out, maybe on base component
|
||||
if self.control_in in socks and socks[self.control_in] == self.zmq.POLLIN:
|
||||
msg = self.control_in.recv()
|
||||
event, payload = CONTROL_UNFRAME(msg)
|
||||
|
||||
# -- Heartbeat --
|
||||
if event == CONTROL_PROTOCOL.HEARTBEAT:
|
||||
# Heart outgoing
|
||||
heartbeat_frame = CONTROL_FRAME(
|
||||
CONTROL_PROTOCOL.OK,
|
||||
payload
|
||||
)
|
||||
self.control_out.send(heartbeat_frame)
|
||||
|
||||
# -- Soft Kill --
|
||||
elif event == CONTROL_PROTOCOL.SHUTDOWN:
|
||||
self.signal_done()
|
||||
self.shutdown()
|
||||
|
||||
# -- Hard Kill --
|
||||
elif event == CONTROL_PROTOCOL.KILL:
|
||||
self.kill()
|
||||
|
||||
|
||||
if self.pull_socket in socks and socks[self.pull_socket] == self.zmq.POLLIN:
|
||||
message = self.pull_socket.recv()
|
||||
|
||||
if message == str(CONTROL_PROTOCOL.DONE):
|
||||
self.ds_finished_counter += 1
|
||||
|
||||
if len(self.data_buffer) == self.ds_finished_counter:
|
||||
#drain any remaining messages in the buffer
|
||||
qutil.LOGGER.debug("draining feed")
|
||||
self.drain()
|
||||
self.signal_done()
|
||||
else:
|
||||
try:
|
||||
event = self.unframe(message)
|
||||
# deserialization error
|
||||
except zp.INVALID_DATASOURCE_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
try:
|
||||
self.append(event)
|
||||
self.send_next()
|
||||
|
||||
# Invalid message
|
||||
except zp.INVALID_DATASOURCE_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
def unframe(self, msg):
|
||||
return zp.DATASOURCE_UNFRAME(msg)
|
||||
|
||||
def frame(self, event):
|
||||
return zp.FEED_FRAME(event)
|
||||
|
||||
# -------------
|
||||
# Flow Control
|
||||
# -------------
|
||||
|
||||
def drain(self):
|
||||
"""
|
||||
Send all messages in the buffer.
|
||||
"""
|
||||
self.draining = True
|
||||
while self.pending_messages() > 0:
|
||||
self.send_next()
|
||||
|
||||
def send_next(self):
|
||||
"""
|
||||
Send the (chronologically) next message in the buffer.
|
||||
"""
|
||||
if not (self.is_full() or self.draining):
|
||||
return
|
||||
|
||||
event = self.next()
|
||||
if(event != None):
|
||||
self.feed_socket.send(self.frame(event), self.zmq.NOBLOCK)
|
||||
self.sent_counters[event.source_id] += 1
|
||||
self.sent_count += 1
|
||||
|
||||
def append(self, event):
|
||||
"""
|
||||
Add an event to the buffer for the source specified by
|
||||
source_id.
|
||||
"""
|
||||
self.data_buffer[event.source_id].append(event)
|
||||
self.recv_counters[event.source_id] += 1
|
||||
self.received_count += 1
|
||||
|
||||
def next(self):
|
||||
"""
|
||||
Get the next message in chronological order.
|
||||
"""
|
||||
if not(self.is_full() or self.draining):
|
||||
return
|
||||
|
||||
cur_source = None
|
||||
earliest_source = None
|
||||
earliest_event = None
|
||||
#iterate over the queues of events from all sources
|
||||
#(1 queue per datasource)
|
||||
for events in self.data_buffer.values():
|
||||
if len(events) == 0:
|
||||
continue
|
||||
cur_source = events
|
||||
first_in_list = events[0]
|
||||
if first_in_list.dt == None:
|
||||
#this is a filler event, discard
|
||||
events.pop(0)
|
||||
continue
|
||||
|
||||
if (earliest_event == None) or (first_in_list.dt <= earliest_event.dt):
|
||||
earliest_event = first_in_list
|
||||
earliest_source = cur_source
|
||||
|
||||
if earliest_event != None:
|
||||
return earliest_source.pop(0)
|
||||
|
||||
def is_full(self):
|
||||
"""
|
||||
Indicates whether the buffer has messages in buffer for
|
||||
all un-DONE, blocking sources.
|
||||
"""
|
||||
for source_id, events in self.data_buffer.iteritems():
|
||||
if len(events) == 0:
|
||||
return False
|
||||
return True
|
||||
|
||||
def pending_messages(self):
|
||||
"""
|
||||
Returns the count of all events from all sources in the
|
||||
buffer.
|
||||
"""
|
||||
total = 0
|
||||
for events in self.data_buffer.values():
|
||||
total += len(events)
|
||||
return total
|
||||
|
||||
def add_source(self, source_id):
|
||||
"""
|
||||
Add a data source to the buffer.
|
||||
"""
|
||||
self.data_buffer[source_id] = []
|
||||
|
||||
def __len__(self):
|
||||
"""
|
||||
Buffer's length is same as internal map holding separate
|
||||
sorted arrays of events keyed by source id.
|
||||
"""
|
||||
return len(self.data_buffer)
|
||||
|
||||
|
||||
class Merge(Feed):
|
||||
"""
|
||||
Merges multiple streams of events into single messages.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
Feed.__init__(self)
|
||||
|
||||
self.init()
|
||||
|
||||
def init(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return "MERGE"
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
return COMPONENT_TYPE.CONDUIT
|
||||
|
||||
def open(self):
|
||||
self.pull_socket = self.bind_merge()
|
||||
self.feed_socket = self.bind_result()
|
||||
|
||||
def next(self):
|
||||
"""Get the next merged message from the feed buffer."""
|
||||
if not (self.is_full() or self.draining):
|
||||
return
|
||||
|
||||
if self.pending_messages() == 0:
|
||||
return
|
||||
|
||||
#
|
||||
#get the raw event from the passthrough transform.
|
||||
result = self.data_buffer[zp.TRANSFORM_TYPE.PASSTHROUGH].pop(0).PASSTHROUGH
|
||||
for source, events in self.data_buffer.iteritems():
|
||||
if source == zp.TRANSFORM_TYPE.PASSTHROUGH:
|
||||
continue
|
||||
if len(events) > 0:
|
||||
cur = events.pop(0)
|
||||
result.merge(cur)
|
||||
return result
|
||||
|
||||
def unframe(self, msg):
|
||||
return zp.TRANSFORM_UNFRAME(msg)
|
||||
|
||||
def frame(self, event):
|
||||
return zp.MERGE_FRAME(event)
|
||||
|
||||
def append(self, event):
|
||||
"""
|
||||
:param event: a namedict with one entry. key is the name of the
|
||||
transform, value is the transformed value.
|
||||
Add an event to the buffer for the source specified by
|
||||
source_id.
|
||||
"""
|
||||
|
||||
self.data_buffer[event.keys()[0]].append(event)
|
||||
self.received_count += 1
|
||||
|
||||
|
||||
class BaseTransform(Component):
|
||||
"""
|
||||
Top level execution entry point for the transform
|
||||
|
||||
- connects to the feed socket to subscribe to events
|
||||
- connects to the result socket (most oftened bound by a TransformsMerge) to PUSH transforms
|
||||
- processes all messages received from feed, until DONE message received
|
||||
- pushes all transforms
|
||||
- sends DONE to result socket, closes all sockets and context
|
||||
|
||||
Parent class for feed transforms. Subclass and override transform
|
||||
method to create a new derived value from the combined feed.
|
||||
"""
|
||||
|
||||
def __init__(self, name, **kwargs):
|
||||
Component.__init__(self)
|
||||
|
||||
self.state = {
|
||||
'name': name
|
||||
}
|
||||
|
||||
self.init(**kwargs)
|
||||
|
||||
def init(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return self.state['name']
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
return COMPONENT_TYPE.CONDUIT
|
||||
|
||||
def open(self):
|
||||
"""
|
||||
Establishes zmq connections.
|
||||
"""
|
||||
#create the feed.
|
||||
self.feed_socket = self.connect_feed()
|
||||
#create the result PUSH
|
||||
self.result_socket = self.connect_merge()
|
||||
|
||||
def do_work(self):
|
||||
"""
|
||||
Loops until feed's DONE message is received:
|
||||
|
||||
- receive an event from the data feed
|
||||
- call transform (subclass' method) on event
|
||||
- send the transformed event
|
||||
|
||||
"""
|
||||
socks = dict(self.poll.poll(self.heartbeat_timeout))
|
||||
|
||||
# TODO: Abstract this out, maybe on base component
|
||||
if self.control_in in socks and socks[self.control_in] == self.zmq.POLLIN:
|
||||
msg = self.control_in.recv()
|
||||
event, payload = CONTROL_UNFRAME(msg)
|
||||
|
||||
# -- Heartbeat --
|
||||
if event == CONTROL_PROTOCOL.HEARTBEAT:
|
||||
# Heart outgoing
|
||||
heartbeat_frame = CONTROL_FRAME(
|
||||
CONTROL_PROTOCOL.OK,
|
||||
payload
|
||||
)
|
||||
self.control_out.send(heartbeat_frame)
|
||||
|
||||
# -- Soft Kill --
|
||||
elif event == CONTROL_PROTOCOL.SHUTDOWN:
|
||||
self.signal_done()
|
||||
self.shutdown()
|
||||
|
||||
# -- Hard Kill --
|
||||
elif event == CONTROL_PROTOCOL.KILL:
|
||||
self.kill()
|
||||
|
||||
if self.feed_socket in socks and socks[self.feed_socket] == self.zmq.POLLIN:
|
||||
message = self.feed_socket.recv()
|
||||
|
||||
if message == str(CONTROL_PROTOCOL.DONE):
|
||||
self.signal_done()
|
||||
return
|
||||
|
||||
try:
|
||||
event = self.unframe(message)
|
||||
except zp.INVALID_FEED_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
try:
|
||||
cur_state = self.transform(event)
|
||||
|
||||
# This is overloaded, so it can fail in all sorts of
|
||||
# unknown ways. Its best to catch it in the
|
||||
# Transformer itself.
|
||||
except Exception as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
try:
|
||||
transform_frame = self.frame(cur_state)
|
||||
except zp.INVALID_TRANSFORM_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
self.result_socket.send(transform_frame, self.zmq.NOBLOCK)
|
||||
|
||||
def frame(self, cur_state):
|
||||
return zp.TRANSFORM_FRAME(cur_state['name'], cur_state['value'])
|
||||
|
||||
def unframe(self, msg):
|
||||
return zp.FEED_UNFRAME(msg)
|
||||
|
||||
def transform(self, event):
|
||||
"""
|
||||
Must return the transformed value as a map with::
|
||||
|
||||
{name:"name of new transform", value: "value of new field"}
|
||||
|
||||
Transforms run in parallel and results are merged into a single map, so
|
||||
transform names must be unique. Best practice is to use the self.state
|
||||
object initialized from the transform configuration, and only set the
|
||||
transformed value::
|
||||
|
||||
self.state['value'] = transformed_value
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class PassthroughTransform(BaseTransform):
|
||||
"""
|
||||
A bypass transform which is also an identity transform::
|
||||
|
||||
+-------+
|
||||
+---| f |--->
|
||||
+-------+
|
||||
+------id------->
|
||||
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs):
|
||||
BaseTransform.__init__(self, "PASSTHROUGH")
|
||||
self.init(**kwargs)
|
||||
|
||||
def init(self, **kwargs):
|
||||
pass
|
||||
|
||||
@property
|
||||
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 {'name':zp.TRANSFORM_TYPE.PASSTHROUGH, 'value': zp.FEED_FRAME(event) }
|
||||
|
||||
|
||||
class DataSource(Component):
|
||||
"""
|
||||
Baseclass for data sources. Subclass and implement send_all - usually this
|
||||
means looping through all records in a store, converting to a dict, and
|
||||
calling send(map).
|
||||
|
||||
Every datasource has a dict property to hold filters::
|
||||
- key -- name of the filter, e.g. SID
|
||||
- value -- a primitive representing the filter. e.g. a list of ints.
|
||||
|
||||
Modify the datasource's filters via the set_filter(name, value)
|
||||
"""
|
||||
def __init__(self, source_id):
|
||||
Component.__init__(self)
|
||||
|
||||
self.id = source_id
|
||||
self.init()
|
||||
self.filter = {}
|
||||
|
||||
def init(self):
|
||||
self.cur_event = None
|
||||
|
||||
def set_filter(self, name, value):
|
||||
self.filter[name] = value
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return self.id
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
return COMPONENT_TYPE.SOURCE
|
||||
|
||||
def open(self):
|
||||
self.data_socket = self.connect_data()
|
||||
|
||||
def send(self, event):
|
||||
"""
|
||||
Emit data.
|
||||
"""
|
||||
assert isinstance(event, zp.namedict)
|
||||
|
||||
event['source_id'] = self.get_id
|
||||
event['type'] = self.get_type
|
||||
|
||||
try:
|
||||
ds_frame = self.frame(event)
|
||||
except zp.INVALID_DATASOURCE_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
self.data_socket.send(ds_frame)
|
||||
|
||||
def frame(self, event):
|
||||
return zp.DATASOURCE_FRAME(event)
|
||||
@@ -1,621 +0,0 @@
|
||||
import time
|
||||
import gevent
|
||||
import itertools
|
||||
# pyzmq
|
||||
import zmq
|
||||
import gevent_zeromq
|
||||
|
||||
from collections import OrderedDict
|
||||
|
||||
from protocol import CONTROL_PROTOCOL, CONTROL_FRAME, \
|
||||
CONTROL_UNFRAME, CONTROL_STATES, INVALID_CONTROL_FRAME \
|
||||
|
||||
states = CONTROL_STATES
|
||||
|
||||
from gpoll import _Poller as GeventPoller
|
||||
|
||||
# Roll Call ( Discovery )
|
||||
# -----------------------
|
||||
#
|
||||
# Controller ( 'foo', 'bar', 'fizz', 'pop' )
|
||||
# ------------------
|
||||
# | | | |
|
||||
# +---+
|
||||
# | 0 | ? ? ?
|
||||
# +---+
|
||||
# |
|
||||
# IDENTITY: foo
|
||||
# get message: PROTOCOL.HEARTBEAT
|
||||
# reply with PROTOCOL.OK
|
||||
#
|
||||
# Controller topology = ( 'foo', 'bar', 'fizz', 'pop' )
|
||||
# 'foo' in topology = YES ->
|
||||
# track 'foo'
|
||||
# ------------------
|
||||
# | | | |
|
||||
# +---+
|
||||
# | 1 | ? ? ?
|
||||
# +---+
|
||||
|
||||
# Heartbeating
|
||||
# ------------
|
||||
#
|
||||
# Controller ( time = 2.717828 )
|
||||
# ------------------
|
||||
# | | | |
|
||||
# +---+ +---+ +---+ +---+
|
||||
# | 0 | | 0 | | 0 | | 0 |
|
||||
# +---+ +---+ +---+ +---+
|
||||
# |
|
||||
# IDENTITY: foo
|
||||
# get message: time = 2.717828
|
||||
# reply with [ foo, 2.71828 ]
|
||||
#
|
||||
# Controller ( foo.status = OK )
|
||||
# ------------------
|
||||
# | | | |
|
||||
# +---+ +---+ +---+ +---+
|
||||
# | 1 | | 0 | | 0 | | 0 |
|
||||
# +---+ +---+ +---+ +---+
|
||||
# |
|
||||
# Controller tracks this node as good
|
||||
# for this heartbeat
|
||||
|
||||
# Shutdown
|
||||
# --------
|
||||
#
|
||||
# Controller ( state = RUNNING )
|
||||
# ------------------
|
||||
# | | | |
|
||||
# +---+ +---+ +---+ +---+
|
||||
# | 1 | | 1 | | 1 | | 1 |
|
||||
# +---+ +---+ +---+ +---+
|
||||
# |
|
||||
# IDENTITY: foo
|
||||
# send [ DONE ]
|
||||
|
||||
# Controller ( state = SHUTDOWN )
|
||||
# Controller topology.remove('foo')
|
||||
# ------------------
|
||||
# | | |
|
||||
# +---+ +---+ +---+ +---+
|
||||
# | | | 1 | | 1 | | 1 |
|
||||
# +---+ +---+ +---+ +---+
|
||||
# |
|
||||
# IDENTITY: foo
|
||||
# yield, stop sending messages
|
||||
|
||||
# Termination
|
||||
# ------------
|
||||
#
|
||||
# Controller ( state = TERMINATE )
|
||||
# ------------------
|
||||
# | | | |
|
||||
# +---+ +---+ +---+ +---+
|
||||
# | 1 | | 1 | | 1 | | 1 |
|
||||
# +---+ +---+ +---+ +---+
|
||||
# |
|
||||
# get message PROTOCOL.KILL
|
||||
|
||||
# Controller ( state = TERMINATE )
|
||||
# ------------------
|
||||
# | | | |
|
||||
# +---+ +---+ +---+ +---+
|
||||
# | 0 | | 0 | | 0 | | 0 |
|
||||
# +---+ +---+ +---+ +---+
|
||||
|
||||
INIT, SOURCES_READY, RUNNING, TERMINATE = CONTROL_STATES
|
||||
|
||||
state_transitions = frozenset([
|
||||
(-1 , INIT),
|
||||
(INIT , SOURCES_READY),
|
||||
(SOURCES_READY , RUNNING),
|
||||
(INIT , TERMINATE),
|
||||
(SOURCES_READY , TERMINATE),
|
||||
(RUNNING , TERMINATE),
|
||||
])
|
||||
|
||||
class UnknownChatter(Exception):
|
||||
def __init__(self, name):
|
||||
self.named = name
|
||||
def __str__(self):
|
||||
return """Component calling itself "%s" talking on unexpected channel"""\
|
||||
% self.named
|
||||
|
||||
class Controller(object):
|
||||
"""
|
||||
A N to M messaging system for inter component communication.
|
||||
|
||||
:param pub_socket: Socket to publish messages, the starting
|
||||
point of :func message_listener: .
|
||||
|
||||
:param route_socket: Socket to listen for status updates for
|
||||
the individual components.
|
||||
:func message_sender: .
|
||||
|
||||
:param logging: Logging interface for tracking broker state
|
||||
Defaults to None
|
||||
|
||||
Topology is the set of components we expect to show up.
|
||||
States are the transitions the sytems go through. The
|
||||
simplest is from RUNNING -> NOT RUNNING .
|
||||
|
||||
Usage::
|
||||
|
||||
controller = Controller(
|
||||
'tcp://127.0.0.1:5000',
|
||||
'tcp://127.0.0.1:5001',
|
||||
)
|
||||
|
||||
# typically you'd want to run this async to your main
|
||||
# program since it blocks indefinetely.
|
||||
controller.manage(
|
||||
[ TOPOLOGY ]
|
||||
[ STATES ]
|
||||
)
|
||||
|
||||
"""
|
||||
|
||||
debug = False
|
||||
period = 1
|
||||
|
||||
def __init__(self, pub_socket, route_socket, logging = None):
|
||||
|
||||
self.context = None
|
||||
self.zmq = None
|
||||
self.zmq_poller = None
|
||||
|
||||
self.running = False
|
||||
self.polling = False
|
||||
self.tracked = set()
|
||||
self.responses = set()
|
||||
|
||||
self.ctime = 0
|
||||
self.tic = time.time()
|
||||
self.freeform = False
|
||||
self._state = -1
|
||||
|
||||
self.associated = []
|
||||
|
||||
self.pub_socket = pub_socket
|
||||
self.route_socket = route_socket
|
||||
|
||||
self.error_replay = OrderedDict()
|
||||
|
||||
if logging:
|
||||
self.logging = logging
|
||||
else:
|
||||
import util as qutil
|
||||
self.logging = qutil.LOGGER
|
||||
|
||||
def init_zmq(self, flavor):
|
||||
|
||||
assert self.zmq_flavor in ['thread', 'mp', 'green']
|
||||
|
||||
if flavor == 'mp':
|
||||
self.zmq = zmq
|
||||
self.context = self.zmq.Context()
|
||||
self.zmq_poller = self.zmq.Poller
|
||||
return
|
||||
if flavor == 'thread':
|
||||
self.zmq = zmq
|
||||
self.context = self.zmq.Context.instance()
|
||||
self.zmq_poller = self.zmq.Poller
|
||||
return
|
||||
if flavor == 'green':
|
||||
self.zmq = gevent_zeromq.zmq
|
||||
self.context = self.zmq.Context.instance()
|
||||
self.zmq_poller = GeventPoller
|
||||
return
|
||||
if flavor == 'pypy':
|
||||
self.zmq = zmq
|
||||
self.context = self.zmq.Context.instance()
|
||||
self.zmq_poller = self.zmq.Poller
|
||||
return
|
||||
|
||||
def manage(self, topology, states=None, context=None):
|
||||
"""
|
||||
Give the controller a set set of components to manage and
|
||||
a set of state transitions for the entire system.
|
||||
"""
|
||||
|
||||
# A freeform topology is where we heartbeat with anything
|
||||
# that shows up.
|
||||
if topology == 'freeform':
|
||||
self.freeform = True
|
||||
self.topology = frozenset([])
|
||||
else:
|
||||
self.freeform = False
|
||||
self.topology = frozenset(topology)
|
||||
|
||||
self.polling = True
|
||||
self.state = CONTROL_STATES.INIT
|
||||
|
||||
@property
|
||||
def state(self):
|
||||
return self._state
|
||||
|
||||
@state.setter
|
||||
def state(self, new):
|
||||
old, self._state = self._state, new
|
||||
|
||||
if (old, new) not in state_transitions:
|
||||
raise RuntimeError("[Controller] Invalid State Transition : %s -> %s" %(old, new))
|
||||
else:
|
||||
self.logging.info("[Controller] State Transition : %s -> %s" %(old, new))
|
||||
|
||||
def run(self):
|
||||
self.running = True
|
||||
self.init_zmq(self.zmq_flavor)
|
||||
|
||||
try:
|
||||
return self._poll() # use a python loop
|
||||
except KeyboardInterrupt:
|
||||
self.logging.info('Shutdown event loop')
|
||||
|
||||
def log_status(self):
|
||||
"""
|
||||
Snapshot of the tracked components at every period.
|
||||
"""
|
||||
#self.logging.info("[Controller] Tracking : %s" % ([c for c in self.tracked],))
|
||||
pass
|
||||
|
||||
def replay_errors(self):
|
||||
"""
|
||||
Replay the errors in the order they were reported to the
|
||||
controller.
|
||||
"""
|
||||
return [ a for a in sorted(self.replay_errors.keys())]
|
||||
|
||||
# -------------
|
||||
# Publications
|
||||
# -------------
|
||||
|
||||
def send_heart(self):
|
||||
if not self.running:
|
||||
return
|
||||
|
||||
heartbeat_frame = CONTROL_FRAME(
|
||||
CONTROL_PROTOCOL.HEARTBEAT,
|
||||
str(self.ctime)
|
||||
)
|
||||
self.pub.send(heartbeat_frame)
|
||||
|
||||
def send_hardkill(self):
|
||||
if not self.running:
|
||||
return
|
||||
|
||||
kill_frame = CONTROL_FRAME(
|
||||
CONTROL_PROTOCOL.KILL,
|
||||
''
|
||||
)
|
||||
self.pub.send(kill_frame)
|
||||
|
||||
def send_softkill(self):
|
||||
if not self.running:
|
||||
return
|
||||
|
||||
soft_frame = CONTROL_FRAME(
|
||||
CONTROL_PROTOCOL.SHUTDOWN,
|
||||
''
|
||||
)
|
||||
self.pub.send(soft_frame)
|
||||
|
||||
# -----------
|
||||
# Event Loops
|
||||
# -----------
|
||||
|
||||
def _poll(self):
|
||||
|
||||
assert self.route_socket
|
||||
assert self.pub_socket
|
||||
assert self.cancel_socket
|
||||
|
||||
# -- Publish --
|
||||
# =============
|
||||
self.pub = self.context.socket(self.zmq.PUB)
|
||||
self.pub.bind(self.pub_socket)
|
||||
|
||||
# -- Cancel --
|
||||
# =============
|
||||
assert isinstance(self.cancel_socket,basestring), self.cancel_socket
|
||||
self.cancel = self.context.socket(self.zmq.REP)
|
||||
self.cancel.connect(self.cancel_socket)
|
||||
|
||||
# -- Router --
|
||||
# =============
|
||||
self.router = self.context.socket(self.zmq.ROUTER)
|
||||
self.router.bind(self.route_socket)
|
||||
|
||||
|
||||
poller = self.zmq.Poller()
|
||||
poller.register(self.router, self.zmq.POLLIN)
|
||||
poller.register(self.cancel, self.zmq.POLLIN)
|
||||
|
||||
self.associated += [self.pub, self.router, self.cancel]
|
||||
|
||||
# TODO: actually do this
|
||||
self.state = CONTROL_STATES.SOURCES_READY
|
||||
|
||||
buffer = []
|
||||
|
||||
for i in itertools.count(0):
|
||||
self.log_status()
|
||||
self.responses = set()
|
||||
|
||||
self.ctime = time.time()
|
||||
self.send_heart()
|
||||
|
||||
while self.polling:
|
||||
# Reset the responses for this cycle
|
||||
|
||||
socks = dict(poller.poll(self.period))
|
||||
tic = time.time()
|
||||
|
||||
if tic - self.ctime > self.period:
|
||||
break
|
||||
|
||||
if self.router in socks and socks[self.router] == self.zmq.POLLIN:
|
||||
rawmessage = self.router.recv()
|
||||
|
||||
if rawmessage:
|
||||
buffer.append(rawmessage)
|
||||
|
||||
try:
|
||||
if not self.router.getsockopt(self.zmq.RCVMORE):
|
||||
self.handle_recv(buffer[:])
|
||||
buffer = []
|
||||
except INVALID_CONTROL_FRAME:
|
||||
self.logging.error('Invalid frame', rawmessage)
|
||||
pass
|
||||
|
||||
if self.cancel in socks and socks[self.cancel] == self.zmq.POLLIN:
|
||||
self.logging.info('[Controller] Received Cancellation')
|
||||
rawmessage = self.cancel.recv()
|
||||
self.shutdown(soft=True)
|
||||
break
|
||||
|
||||
self.beat()
|
||||
|
||||
if self.zmq_flavor == 'green':
|
||||
gevent.sleep(0)
|
||||
|
||||
if self.state is CONTROL_STATES.TERMINATE:
|
||||
break
|
||||
|
||||
if not self.polling:
|
||||
break
|
||||
|
||||
# After loop exits
|
||||
self.terminated = True
|
||||
|
||||
def beat(self):
|
||||
|
||||
# These the set overloaded operations
|
||||
# A & B ~ set.intersection
|
||||
# A - B ~ set.difference
|
||||
|
||||
# * good - Components we are currently tracking and who just sent
|
||||
# us back the right response.
|
||||
# * bad - Components we are currently tracking but who did not
|
||||
# send us back a response.
|
||||
# * new - Components we haven't heard from yet, but sent back the
|
||||
# right response.
|
||||
|
||||
good = self.tracked & self.responses
|
||||
bad = self.tracked - good
|
||||
new = self.responses - good
|
||||
|
||||
for component in new:
|
||||
self.new(component)
|
||||
|
||||
for component in bad:
|
||||
self.fail(component)
|
||||
|
||||
# --------------
|
||||
# Init Handlers
|
||||
# --------------
|
||||
|
||||
def new_source(self):
|
||||
if self.state is CONTROL_STATES.RUNNING:
|
||||
self.state = SOURCES_READY
|
||||
|
||||
def new_universal(self):
|
||||
pass
|
||||
|
||||
# The various "states of being that a component can inform us
|
||||
# of
|
||||
def new(self, component):
|
||||
if self.state is CONTROL_STATES.TERMINATE:
|
||||
return
|
||||
|
||||
self.logging.info('[Controller] Now Tracking "%s" ' % component)
|
||||
|
||||
universal = self.new_universal
|
||||
init_handlers = {
|
||||
'FEED' : self.new_source,
|
||||
}
|
||||
|
||||
if component in self.topology or self.freeform:
|
||||
init_handlers.get(component, universal)()
|
||||
self.tracked.add(component)
|
||||
else:
|
||||
# Some sort of socket collision has occured, this is
|
||||
# a very bad failure mode.
|
||||
raise UnknownChatter(component)
|
||||
|
||||
# ------------------
|
||||
# Epic Fail Handling
|
||||
# ------------------
|
||||
|
||||
def fail_universal(self):
|
||||
pass
|
||||
# TODO: this requires higher order functionality
|
||||
#self.logging.error('[Controller] System in exception state, shutting down')
|
||||
#self.shutdown(soft=True)
|
||||
|
||||
def fail(self, component):
|
||||
if self.state is CONTROL_STATES.TERMINATE:
|
||||
return
|
||||
|
||||
universal = self.fail_universal
|
||||
fail_handlers = { }
|
||||
|
||||
if component in self.topology or self.freeform:
|
||||
self.logging.info('[Controller] Component "%s" timed out' % component)
|
||||
self.tracked.remove(component)
|
||||
fail_handlers.get(component, universal)()
|
||||
|
||||
# -------------------
|
||||
# Completion Handling
|
||||
# -------------------
|
||||
|
||||
def done(self, component):
|
||||
self.logging.info('[Controller] Component "%s" done.' % component)
|
||||
|
||||
# --------------
|
||||
# Error Handling
|
||||
# --------------
|
||||
|
||||
def exception_universal(self):
|
||||
"""
|
||||
Shutdown the system on failure.
|
||||
"""
|
||||
self.logging.error('[Controller] System in exception state, shutting down')
|
||||
self.shutdown(soft=True)
|
||||
|
||||
def exception(self, component, failure):
|
||||
universal = self.exception_universal
|
||||
exception_handlers = { }
|
||||
|
||||
if component in self.topology or self.freeform:
|
||||
self.error_replay[(component, time.time())] = failure
|
||||
self.logging.error('[Controller] Component "%s" in exception state' % component)
|
||||
|
||||
exception_handlers.get(component, universal)()
|
||||
else:
|
||||
raise UnknownChatter(component)
|
||||
|
||||
# -----------------
|
||||
# Protocol Handling
|
||||
# -----------------
|
||||
|
||||
def handle_recv(self, msg):
|
||||
"""
|
||||
Check for proper framing at the transport layer.
|
||||
Seperates the proper frames from anything else that might
|
||||
be coming over the wire. Which shouldn't happen ... right?
|
||||
"""
|
||||
identity = msg[0]
|
||||
id, status = CONTROL_UNFRAME(msg[1])
|
||||
|
||||
# A component is telling us its alive:
|
||||
if id is CONTROL_PROTOCOL.OK:
|
||||
|
||||
if status == str(self.ctime):
|
||||
self.responses.add(identity)
|
||||
else:
|
||||
# Otherwise its something weird and we don't know
|
||||
# what to do so just say so
|
||||
self.logging.error("Weird stuff happened: %s" % msg)
|
||||
|
||||
# A component is telling us it failed, and how
|
||||
if id is CONTROL_PROTOCOL.EXCEPTION:
|
||||
self.exception(identity, status)
|
||||
|
||||
# A component is telling us its done with work and won't
|
||||
# be talking to us anymore
|
||||
if id is CONTROL_PROTOCOL.DONE:
|
||||
self.done(identity)
|
||||
|
||||
# -------------------
|
||||
# Hooks for Endpoints
|
||||
# -------------------
|
||||
|
||||
# These are all connects so no complex allocation logic is
|
||||
# needed. Dealers and Subscribers can all come and go as a
|
||||
# function of time without impacting flow of the whole
|
||||
# system.
|
||||
|
||||
def message_sender(self, identity, context = None):
|
||||
"""
|
||||
Spin off a socket used for sending messages to this
|
||||
controller.
|
||||
"""
|
||||
|
||||
if not context:
|
||||
context = self.zmq.Context.instance()
|
||||
|
||||
s = context.socket(zmq.DEALER)
|
||||
s.setsockopt(zmq.IDENTITY, identity)
|
||||
s.connect(self.route_socket)
|
||||
|
||||
self.associated.append(s)
|
||||
return s
|
||||
|
||||
def message_listener(self, context = None):
|
||||
"""
|
||||
Spin off a socket used for receiving messages from this
|
||||
controller.
|
||||
"""
|
||||
|
||||
if not context:
|
||||
context = self.zmq.Context.instance()
|
||||
|
||||
s = context.socket(zmq.SUB)
|
||||
s.connect(self.pub_socket)
|
||||
s.setsockopt(zmq.SUBSCRIBE, '')
|
||||
|
||||
self.associated.append(s)
|
||||
return s
|
||||
|
||||
def do_error_replay(self):
|
||||
for (component, time), error in self.error_replay.iteritems():
|
||||
self.logging.info('[Controller] Error Log for -- %s --:\n%s' %
|
||||
(component, error))
|
||||
|
||||
def shutdown(self, hard=False, soft=True, context=None):
|
||||
|
||||
if not self.polling:
|
||||
return
|
||||
|
||||
self.polling = False
|
||||
|
||||
assert hard or soft, """ Must specify kill hard or soft """
|
||||
|
||||
if hard:
|
||||
self.state = CONTROL_STATES.TERMINATE
|
||||
|
||||
self.logging.info('[Controller] Hard Shutdown')
|
||||
|
||||
#for asoc in self.associated:
|
||||
#asoc.close()
|
||||
|
||||
if soft:
|
||||
self.state = CONTROL_STATES.TERMINATE
|
||||
|
||||
self.logging.info('[Controller] Soft Shutdown')
|
||||
self.send_softkill()
|
||||
|
||||
#for asoc in self.associated:
|
||||
#asoc.close()
|
||||
|
||||
self.do_error_replay()
|
||||
|
||||
if __name__ == '__main__':
|
||||
|
||||
print 'Running on '\
|
||||
'tcp://127.0.0.1:5000 '\
|
||||
'tcp://127.0.0.1:5001 '
|
||||
|
||||
controller = Controller(
|
||||
'tcp://127.0.0.1:5000',
|
||||
'tcp://127.0.0.1:5001',
|
||||
)
|
||||
controller.zmq_flavor = 'green'
|
||||
|
||||
controller.manage(
|
||||
'freeform',
|
||||
[]
|
||||
)
|
||||
controller.run()
|
||||
@@ -0,0 +1,3 @@
|
||||
"""
|
||||
Thomas's parameter optimization library.
|
||||
"""
|
||||
+39
-41
@@ -65,7 +65,7 @@ Namedict
|
||||
Namedicts are dict like objects that have fields accessible by attribute lookup
|
||||
as well as being indexable and iterable::
|
||||
|
||||
HEARTBEAT_PROTOCOL = namedict({
|
||||
HEARTBEAT_PROTOCOL = ndict({
|
||||
'REQ' : b'\x01',
|
||||
'REP' : b'\x02',
|
||||
})
|
||||
@@ -118,13 +118,10 @@ import msgpack
|
||||
import numbers
|
||||
import datetime
|
||||
import pytz
|
||||
import numpy
|
||||
import time
|
||||
import copy
|
||||
from collections import namedtuple
|
||||
|
||||
from protocol_utils import Enum, FrameExceptionFactory, namedict
|
||||
from date_utils import EPOCH, UN_EPOCH
|
||||
from utils.protocol_utils import Enum, FrameExceptionFactory, ndict
|
||||
from utils.date_utils import EPOCH, UN_EPOCH
|
||||
|
||||
# -----------------------
|
||||
# Control Protocol
|
||||
@@ -221,7 +218,7 @@ def DATASOURCE_FRAME(event):
|
||||
Wraps any datasource payload with id and type, so that unpacking may choose
|
||||
the write UNFRAME for the payload.
|
||||
|
||||
:param event: namedict with following properties
|
||||
:param event: ndict with following properties
|
||||
|
||||
- *ds_id* an identifier that is unique to the datasource in the context of a component host (e.g. Simulator)
|
||||
- *ds_type* a string denoting the datasource type. Must be on of:
|
||||
@@ -260,42 +257,43 @@ def DATASOURCE_FRAME(event):
|
||||
|
||||
def DATASOURCE_UNFRAME(msg):
|
||||
"""
|
||||
|
||||
Extracts payload, and calls correct UNFRAME method based on the \
|
||||
datasource type passed along.
|
||||
|
||||
|
||||
Extracts payload, and calls correct UNFRAME method based on the
|
||||
datasource type passed along.
|
||||
|
||||
Returns a dict containing at least:
|
||||
|
||||
|
||||
- source_id
|
||||
- type
|
||||
|
||||
other properties are added based on the datasource type:
|
||||
|
||||
|
||||
- TRADE
|
||||
|
||||
|
||||
- sid - int security identifier
|
||||
- price - float
|
||||
- volume - int
|
||||
- dt - a datetime object
|
||||
|
||||
"""
|
||||
|
||||
"""
|
||||
try:
|
||||
ds_type, source_id, payload = msgpack.loads(msg)
|
||||
assert isinstance(ds_type, int)
|
||||
rval = namedict({'source_id':source_id})
|
||||
|
||||
rval = ndict({'source_id':source_id})
|
||||
|
||||
if payload == DATASOURCE_TYPE.EMPTY:
|
||||
child_value = namedict({'dt':None})
|
||||
child_value = ndict({'dt':None})
|
||||
elif(ds_type == DATASOURCE_TYPE.TRADE):
|
||||
child_value = TRADE_UNFRAME(payload)
|
||||
elif(ds_type == DATASOURCE_TYPE.ORDER):
|
||||
child_value = ORDER_SOURCE_UNFRAME(payload)
|
||||
else:
|
||||
raise INVALID_DATASOURCE_FRAME(msg)
|
||||
|
||||
|
||||
rval.merge(child_value)
|
||||
return rval
|
||||
|
||||
|
||||
except TypeError:
|
||||
raise INVALID_DATASOURCE_FRAME(msg)
|
||||
except ValueError:
|
||||
@@ -309,12 +307,12 @@ INVALID_FEED_FRAME = FrameExceptionFactory('FEED')
|
||||
|
||||
def FEED_FRAME(event):
|
||||
"""
|
||||
:param event: a nameddict with at least
|
||||
|
||||
:param event: a ndict with at least
|
||||
|
||||
- source_id
|
||||
- type
|
||||
"""
|
||||
assert isinstance(event, namedict)
|
||||
assert isinstance(event, ndict)
|
||||
source_id = event.source_id
|
||||
ds_type = event.type
|
||||
PACK_DATE(event)
|
||||
@@ -326,7 +324,7 @@ def FEED_UNFRAME(msg):
|
||||
payload = msgpack.loads(msg)
|
||||
#TODO: anything we can do to assert more about the content of the dict?
|
||||
assert isinstance(payload, dict)
|
||||
rval = namedict(payload)
|
||||
rval = ndict(payload)
|
||||
UNPACK_DATE(rval)
|
||||
return rval
|
||||
except TypeError:
|
||||
@@ -350,13 +348,13 @@ def TRANSFORM_FRAME(name, value):
|
||||
|
||||
def TRANSFORM_UNFRAME(msg):
|
||||
"""
|
||||
:rtype: namedict with <transform_name>:<transform_value>
|
||||
:rtype: ndict with <transform_name>:<transform_value>
|
||||
"""
|
||||
try:
|
||||
|
||||
name, value = msgpack.loads(msg)
|
||||
if(value == TRANSFORM_TYPE.EMPTY):
|
||||
return namedict({name : None})
|
||||
return ndict({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):
|
||||
@@ -364,7 +362,7 @@ def TRANSFORM_UNFRAME(msg):
|
||||
elif(name == TRANSFORM_TYPE.TRANSACTION):
|
||||
value = TRANSACTION_UNFRAME(value)
|
||||
|
||||
return namedict({name : value})
|
||||
return ndict({name : value})
|
||||
except TypeError:
|
||||
raise INVALID_TRANSFORM_FRAME(msg)
|
||||
except ValueError:
|
||||
@@ -382,7 +380,7 @@ def MERGE_FRAME(event):
|
||||
- source_id
|
||||
- type
|
||||
"""
|
||||
assert isinstance(event, namedict)
|
||||
assert isinstance(event, ndict)
|
||||
PACK_DATE(event)
|
||||
if(event.has_attr(TRANSFORM_TYPE.TRANSACTION)):
|
||||
if(event.TRANSACTION == None):
|
||||
@@ -397,7 +395,7 @@ def MERGE_UNFRAME(msg):
|
||||
payload = msgpack.loads(msg)
|
||||
#TODO: anything we can do to assert more about the content of the dict?
|
||||
assert isinstance(payload, dict)
|
||||
payload = namedict(payload)
|
||||
payload = ndict(payload)
|
||||
if(payload.has_attr(TRANSFORM_TYPE.TRANSACTION)):
|
||||
if(payload.TRANSACTION == TRANSFORM_TYPE.EMPTY):
|
||||
payload.TRANSACTION = None
|
||||
@@ -425,7 +423,7 @@ INVALID_TRADE_FRAME = FrameExceptionFactory('TRADE')
|
||||
|
||||
def TRADE_FRAME(event):
|
||||
"""
|
||||
:param event: should be a namedict with:
|
||||
:param event: should be a ndict with:
|
||||
|
||||
- ds_id -- the datasource id sending this trade out
|
||||
- sid -- the security id
|
||||
@@ -434,7 +432,7 @@ def TRADE_FRAME(event):
|
||||
- dt -- datetime for the trade
|
||||
|
||||
"""
|
||||
assert isinstance(event, namedict)
|
||||
assert isinstance(event, ndict)
|
||||
assert event.type == DATASOURCE_TYPE.TRADE
|
||||
assert isinstance(event.sid, int)
|
||||
assert isinstance(event.price, numbers.Real)
|
||||
@@ -456,7 +454,7 @@ def TRADE_UNFRAME(msg):
|
||||
assert isinstance(sid, int)
|
||||
assert isinstance(price, numbers.Real)
|
||||
assert isinstance(volume, numbers.Integral)
|
||||
rval = namedict({
|
||||
rval = ndict({
|
||||
'sid' : sid,
|
||||
'price' : price,
|
||||
'volume' : volume,
|
||||
@@ -491,7 +489,7 @@ def ORDER_UNFRAME(msg):
|
||||
sid, amount, dt = msgpack.loads(msg)
|
||||
assert isinstance(sid, int)
|
||||
assert isinstance(amount, int)
|
||||
rval = namedict({
|
||||
rval = ndict({
|
||||
'sid':sid,
|
||||
'amount':amount,
|
||||
'dt':dt
|
||||
@@ -513,7 +511,7 @@ def ORDER_UNFRAME(msg):
|
||||
|
||||
|
||||
def TRANSACTION_FRAME(event):
|
||||
assert isinstance(event, namedict)
|
||||
assert isinstance(event, ndict)
|
||||
assert isinstance(event.sid, int)
|
||||
assert isinstance(event.price, numbers.Real)
|
||||
assert isinstance(event.commission, numbers.Real)
|
||||
@@ -535,7 +533,7 @@ def TRANSACTION_UNFRAME(msg):
|
||||
assert isinstance(price, numbers.Real)
|
||||
assert isinstance(commission, numbers.Real)
|
||||
assert isinstance(amount, int)
|
||||
rval = namedict({
|
||||
rval = ndict({
|
||||
'sid' : sid,
|
||||
'price' : price,
|
||||
'amount' : amount,
|
||||
@@ -577,7 +575,7 @@ def ORDER_SOURCE_FRAME(event):
|
||||
def ORDER_SOURCE_UNFRAME(msg):
|
||||
try:
|
||||
sid, amount, dt, source_id, source_type = msgpack.loads(msg)
|
||||
event = namedict({
|
||||
event = ndict({
|
||||
"sid" : sid,
|
||||
"amount" : amount,
|
||||
"dt" : dt,
|
||||
@@ -688,7 +686,7 @@ def PACK_DATE(event):
|
||||
|
||||
PACK_DATE and UNPACK_DATE are inverse operations.
|
||||
|
||||
:param event: event must a namedict with a property named 'dt' that is a datetime.
|
||||
:param event: event must a ndict with a property named 'dt' that is a datetime.
|
||||
:rtype: None
|
||||
"""
|
||||
assert isinstance(event.dt, datetime.datetime)
|
||||
@@ -710,7 +708,7 @@ def UNPACK_DATE(event):
|
||||
|
||||
UNPACK_DATE and PACK_DATE are inverse operations.
|
||||
|
||||
:param tuple event: event must a namedict with:
|
||||
:param tuple event: event must a ndict with:
|
||||
|
||||
- a property named 'dt_tuple' that is a tuple of integers \
|
||||
representing the date and time in UTC.
|
||||
@@ -742,15 +740,15 @@ ORDER_PROTOCOL = Enum(
|
||||
)
|
||||
|
||||
|
||||
#Transform type needs to be a namedict to facilitate merging.
|
||||
TRANSFORM_TYPE = namedict({
|
||||
#Transform type needs to be a ndict to facilitate merging.
|
||||
TRANSFORM_TYPE = ndict({
|
||||
'TRANSACTION' : 'TRANSACTION', #needed?
|
||||
'PASSTHROUGH' : 'PASSTHROUGH',
|
||||
'EMPTY' : ''
|
||||
})
|
||||
|
||||
|
||||
FINANCE_COMPONENT = namedict({
|
||||
FINANCE_COMPONENT = ndict({
|
||||
'TRADING_CLIENT' : 'TRADING_CLIENT',
|
||||
'PORTFOLIO_CLIENT' : 'PORTFOLIO_CLIENT',
|
||||
'ORDER_SOURCE' : 'ORDER_SOURCE',
|
||||
|
||||
@@ -1,221 +0,0 @@
|
||||
import copy
|
||||
import pandas
|
||||
from ctypes import Structure, c_ubyte
|
||||
from collections import MutableMapping
|
||||
from itertools import izip
|
||||
|
||||
def Enum(*options):
|
||||
"""
|
||||
Fast enums are very important when we want really tight zmq
|
||||
loops. These are probably going to evolve into pure C structs
|
||||
anyways so might as well get going on that.
|
||||
"""
|
||||
class cstruct(Structure):
|
||||
_fields_ = [(o, c_ubyte) for o in options]
|
||||
__iter__ = lambda s: iter(range(len(options)))
|
||||
return cstruct(*range(len(options)))
|
||||
|
||||
def FrameExceptionFactory(name):
|
||||
"""
|
||||
Exception factory with a closure around the frame class name.
|
||||
"""
|
||||
class InvalidFrame(Exception):
|
||||
def __init__(self, got):
|
||||
self.got = got
|
||||
|
||||
def __str__(self):
|
||||
return "Invalid {framecls} Frame: {got}".format(
|
||||
framecls = name,
|
||||
got = self.got,
|
||||
)
|
||||
|
||||
return InvalidFrame
|
||||
|
||||
class namedict(MutableMapping):
|
||||
"""
|
||||
|
||||
Namedicts are dict like objects that have fields accessible by attribute lookup
|
||||
as well as being indexable and iterable::
|
||||
|
||||
HEARTBEAT_PROTOCOL = namedict({
|
||||
'REQ' : b'\x01',
|
||||
'REP' : b'\x02',
|
||||
})
|
||||
|
||||
HEARTBEAT_PROTOCOL.REQ # syntactic sugar
|
||||
HEARTBEAT_PROTOCOL.REP # oh suga suga
|
||||
|
||||
For more complex structs use collections.namedtuple:
|
||||
"""
|
||||
|
||||
def __init__(self, dct=None):
|
||||
if(dct):
|
||||
self.__dict__.update(dct)
|
||||
|
||||
def __setitem__(self, key, value):
|
||||
"""
|
||||
Required for use by pymongo as_class parameter to find.
|
||||
"""
|
||||
if(key == '_id'):
|
||||
self.__dict__['id'] = value
|
||||
else:
|
||||
self.__dict__[key] = value
|
||||
|
||||
def __getitem__(self, key):
|
||||
return self.__dict__[key]
|
||||
|
||||
def __delitem__(self, key):
|
||||
del self.__dict__[key]
|
||||
|
||||
def __iter__(self):
|
||||
return self.__dict__.iterkeys()
|
||||
|
||||
def __len__(self):
|
||||
return len(self.__dict__)
|
||||
|
||||
def keys(self):
|
||||
return self.__dict__.keys()
|
||||
|
||||
def as_dict(self):
|
||||
# shallow copy is O(n)
|
||||
return copy.copy(self.__dict__)
|
||||
|
||||
def delete(self, key):
|
||||
del(self.__dict__[key])
|
||||
|
||||
def merge(self, other_nd):
|
||||
assert isinstance(other_nd, namedict)
|
||||
self.__dict__.update(other_nd.__dict__)
|
||||
|
||||
def __repr__(self):
|
||||
return "namedict: " + str(self.__dict__)
|
||||
|
||||
def __eq__(self, other):
|
||||
# !!!!!!!!!!!!!!!!!!!!
|
||||
# !!!! DANGEROUS !!!!!
|
||||
# !!!!!!!!!!!!!!!!!!!!
|
||||
return other != None and self.__dict__ == other.__dict__
|
||||
|
||||
def has_attr(self, name):
|
||||
return self.__dict__.has_key(name)
|
||||
|
||||
def as_series(self):
|
||||
s = pandas.Series(self.__dict__)
|
||||
s.name = self.sid
|
||||
return s
|
||||
|
||||
class ndict(MutableMapping):
|
||||
"""
|
||||
Xtreme Namedicts 2.0
|
||||
|
||||
Ndicts are dict like objects that have fields accessible by attribute
|
||||
lookup as well as being indexable and iterable. Done right
|
||||
this time.
|
||||
"""
|
||||
|
||||
def __init__(self, dct=None):
|
||||
self.__internal = dict()
|
||||
self.cls = frozenset(dir(self))
|
||||
|
||||
if dct:
|
||||
self.__internal.update(dct)
|
||||
|
||||
# Abstact Overloads
|
||||
# -----------------
|
||||
|
||||
def __setitem__(self, key, value):
|
||||
"""
|
||||
Required for use by pymongo as_class parameter to find.
|
||||
"""
|
||||
if key == '_id':
|
||||
self.__internal['id'] = value
|
||||
else:
|
||||
self.__internal[key] = value
|
||||
|
||||
|
||||
def __getattr__(self, key):
|
||||
if key in self.cls:
|
||||
return self.__dict__[key]
|
||||
else:
|
||||
return self.__internal[key]
|
||||
|
||||
def __getitem__(self, key):
|
||||
return self.__internal[key]
|
||||
|
||||
def __delitem__(self, key):
|
||||
del self.__internal[key]
|
||||
|
||||
def __iter__(self):
|
||||
return self.__internal.iterkeys()
|
||||
|
||||
def __len__(self):
|
||||
return len(self.__internal)
|
||||
|
||||
# Compatability with namedicts
|
||||
# ----------------------------
|
||||
|
||||
# for compat, not the Python way to do things though...
|
||||
# Deprecated, use builtin ``del`` operator.
|
||||
delete = __delitem__
|
||||
|
||||
def has_attr(self, key):
|
||||
"""
|
||||
Deprecated, use builtin ``in`` operator.
|
||||
"""
|
||||
return self.__contains__(key)
|
||||
|
||||
def has_key(self, key):
|
||||
return self.__contains__(key)
|
||||
|
||||
# Custom Methods
|
||||
# --------------
|
||||
|
||||
def copy(self):
|
||||
return ndict(copy.copy(self.__internal))
|
||||
|
||||
def as_dataframe(self):
|
||||
"""
|
||||
Return the representation as a Pandas dataframe.
|
||||
"""
|
||||
d = pandas.DataFrame(self.__internal)
|
||||
return d
|
||||
|
||||
def as_series(self):
|
||||
"""
|
||||
Return the representation as a Pandas time series.
|
||||
"""
|
||||
s = pandas.Series(self.__internal)
|
||||
s.name = self.sid
|
||||
return s
|
||||
|
||||
def as_dict(self):
|
||||
"""
|
||||
Return the representation as a vanilla Python dict.
|
||||
"""
|
||||
# shallow copy is O(n)
|
||||
return copy.copy(self.__internal)
|
||||
|
||||
def merge(self, other_nd):
|
||||
"""
|
||||
Merge in place with another ndict.
|
||||
"""
|
||||
assert isinstance(other_nd, ndict)
|
||||
self.__internal.update(other_nd.__internal)
|
||||
|
||||
def __repr__(self):
|
||||
return "namedict: " + str(self.__internal)
|
||||
|
||||
# Faster dictionary comparison?
|
||||
#def __eq__(self, other):
|
||||
#assert isinstance(other, ndict)
|
||||
|
||||
#keyeq = set(self.keys()) == set(other.keys())
|
||||
|
||||
#if not keyeq:
|
||||
#return False
|
||||
|
||||
#for i, j in izip(self.itervalues(), other.itervalues()):
|
||||
#if i != j:
|
||||
#return False
|
||||
|
||||
#return True
|
||||
@@ -1,22 +0,0 @@
|
||||
from zipline.messaging import BaseTransform
|
||||
from zipline.protocol import COMPONENT_TYPE
|
||||
|
||||
class DivideByZeroTransform(BaseTransform):
|
||||
"""
|
||||
A transform that fails.
|
||||
"""
|
||||
|
||||
def __init__(self, name):
|
||||
BaseTransform.__init__(self, "PASSTHROUGH")
|
||||
self.state['name'] = name
|
||||
self.init()
|
||||
|
||||
def init(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
return COMPONENT_TYPE.CONDUIT
|
||||
|
||||
def transform(self, event):
|
||||
return { 'value': 0/0 }
|
||||
@@ -6,14 +6,20 @@ Transforms provide re-useable components for stream processing. All
|
||||
Transforms expect to receive data events from zipline.core.DataFeed
|
||||
asynchronously via zeromq. Each transform is designed to run in independent
|
||||
process space, independently of all other transforms, to allow for parallel
|
||||
computation.
|
||||
computation.
|
||||
|
||||
Each transform must maintain the state necessary to calculate the transform of
|
||||
each new feed events.
|
||||
each new feed events.
|
||||
|
||||
To simplify the consumption of feed and transform data events, this module
|
||||
also provides the TransformsMerge class. TransformsMerge initializes as set of
|
||||
transforms and subscribes to their output. Each feed event is then combined with
|
||||
all the transforms of that event into a single new message.
|
||||
|
||||
"""
|
||||
"""
|
||||
|
||||
from base import BaseTransform
|
||||
|
||||
__all__ = [
|
||||
BaseTransform,
|
||||
]
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
import logging
|
||||
from zipline.core.component import Component
|
||||
|
||||
import zipline.protocol as zp
|
||||
from zipline.protocol import CONTROL_PROTOCOL, COMPONENT_TYPE, \
|
||||
CONTROL_FRAME, CONTROL_UNFRAME
|
||||
|
||||
LOGGER = logging.getLogger('ZiplineLogger')
|
||||
|
||||
class BaseTransform(Component):
|
||||
"""
|
||||
Top level execution entry point for the transform
|
||||
|
||||
- connects to the feed socket to subscribe to events
|
||||
- connects to the result socket (most oftened bound by a TransformsMerge) to PUSH transforms
|
||||
- processes all messages received from feed, until DONE message received
|
||||
- pushes all transforms
|
||||
- sends DONE to result socket, closes all sockets and context
|
||||
|
||||
Parent class for feed transforms. Subclass and override transform
|
||||
method to create a new derived value from the combined feed.
|
||||
"""
|
||||
|
||||
def __init__(self, name, **kwargs):
|
||||
Component.__init__(self)
|
||||
|
||||
self.state = {
|
||||
'name': name
|
||||
}
|
||||
|
||||
self.init(**kwargs)
|
||||
|
||||
def init(self):
|
||||
pass
|
||||
|
||||
@property
|
||||
def get_id(self):
|
||||
return self.state['name']
|
||||
|
||||
@property
|
||||
def get_type(self):
|
||||
return COMPONENT_TYPE.CONDUIT
|
||||
|
||||
def open(self):
|
||||
"""
|
||||
Establishes zmq connections.
|
||||
"""
|
||||
#create the feed.
|
||||
self.feed_socket = self.connect_feed()
|
||||
#create the result PUSH
|
||||
self.result_socket = self.connect_merge()
|
||||
|
||||
def do_work(self):
|
||||
"""
|
||||
Loops until feed's DONE message is received:
|
||||
|
||||
- receive an event from the data feed
|
||||
- call transform (subclass' method) on event
|
||||
- send the transformed event
|
||||
|
||||
"""
|
||||
socks = dict(self.poll.poll(self.heartbeat_timeout))
|
||||
|
||||
# TODO: Abstract this out, maybe on base component
|
||||
if self.control_in in socks and socks[self.control_in] == self.zmq.POLLIN:
|
||||
msg = self.control_in.recv()
|
||||
event, payload = CONTROL_UNFRAME(msg)
|
||||
|
||||
# -- Heartbeat --
|
||||
if event == CONTROL_PROTOCOL.HEARTBEAT:
|
||||
# Heart outgoing
|
||||
heartbeat_frame = CONTROL_FRAME(
|
||||
CONTROL_PROTOCOL.OK,
|
||||
payload
|
||||
)
|
||||
self.control_out.send(heartbeat_frame)
|
||||
|
||||
# -- Soft Kill --
|
||||
elif event == CONTROL_PROTOCOL.SHUTDOWN:
|
||||
self.signal_done()
|
||||
self.shutdown()
|
||||
|
||||
# -- Hard Kill --
|
||||
elif event == CONTROL_PROTOCOL.KILL:
|
||||
self.kill()
|
||||
|
||||
if self.feed_socket in socks and socks[self.feed_socket] == self.zmq.POLLIN:
|
||||
message = self.feed_socket.recv()
|
||||
|
||||
if message == str(CONTROL_PROTOCOL.DONE):
|
||||
self.signal_done()
|
||||
return
|
||||
|
||||
try:
|
||||
event = self.unframe(message)
|
||||
except zp.INVALID_FEED_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
try:
|
||||
cur_state = self.transform(event)
|
||||
|
||||
# This is overloaded, so it can fail in all sorts of
|
||||
# unknown ways. Its best to catch it in the
|
||||
# Transformer itself.
|
||||
except Exception as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
try:
|
||||
transform_frame = self.frame(cur_state)
|
||||
except zp.INVALID_TRANSFORM_FRAME as exc:
|
||||
return self.signal_exception(exc)
|
||||
|
||||
self.result_socket.send(transform_frame, self.zmq.NOBLOCK)
|
||||
|
||||
def frame(self, cur_state):
|
||||
return zp.TRANSFORM_FRAME(cur_state['name'], cur_state['value'])
|
||||
|
||||
def unframe(self, msg):
|
||||
return zp.FEED_UNFRAME(msg)
|
||||
|
||||
def transform(self, event):
|
||||
"""
|
||||
Must return the transformed value as a map with::
|
||||
|
||||
{name:"name of new transform", value: "value of new field"}
|
||||
|
||||
Transforms run in parallel and results are merged into a single map, so
|
||||
transform names must be unique. Best practice is to use the self.state
|
||||
object initialized from the transform configuration, and only set the
|
||||
transformed value::
|
||||
|
||||
self.state['value'] = transformed_value
|
||||
"""
|
||||
raise NotImplementedError
|
||||
@@ -0,0 +1,5 @@
|
||||
from protocol_utils import ndict
|
||||
|
||||
__all__ = [
|
||||
ndict,
|
||||
]
|
||||
@@ -1,19 +1,19 @@
|
||||
"""
|
||||
Factory functions to prepare useful data for tests.
|
||||
"""
|
||||
|
||||
import pytz
|
||||
import msgpack
|
||||
import random
|
||||
|
||||
from datetime import datetime, timedelta
|
||||
import zipline.util as qutil
|
||||
import zipline.finance.risk as risk
|
||||
import zipline.protocol as zp
|
||||
from zipline.sources import SpecificEquityTrades, RandomEquityTrades
|
||||
from zipline.finance.sources import SpecificEquityTrades, RandomEquityTrades
|
||||
from zipline.finance.trading import TradingEnvironment
|
||||
|
||||
def load_market_data():
|
||||
fp_bm = open("./zipline/test/benchmark.msgpack", "rb")
|
||||
fp_bm = open("./tests/benchmark.msgpack", "rb")
|
||||
bm_list = msgpack.loads(fp_bm.read())
|
||||
bm_returns = []
|
||||
for packed_date, returns in bm_list:
|
||||
@@ -24,20 +24,20 @@ def load_market_data():
|
||||
# second=0,
|
||||
# tzinfo=pytz.utc
|
||||
#)
|
||||
|
||||
|
||||
daily_return = risk.DailyReturn(date=event_dt, returns=returns)
|
||||
bm_returns.append(daily_return)
|
||||
bm_returns = sorted(bm_returns, key=lambda(x): x.date)
|
||||
fp_tr = open("./zipline/test/treasury_curves.msgpack", "rb")
|
||||
bm_returns = sorted(bm_returns, key=lambda(x): x.date)
|
||||
fp_tr = open(".//tests/treasury_curves.msgpack", "rb")
|
||||
tr_list = msgpack.loads(fp_tr.read())
|
||||
tr_curves = {}
|
||||
for packed_date, curve in tr_list:
|
||||
tr_dt = zp.tuple_to_date(packed_date)
|
||||
#tr_dt = tr_dt.replace(hour=0, minute=0, second=0, tzinfo=pytz.utc)
|
||||
tr_curves[tr_dt] = curve
|
||||
|
||||
|
||||
return bm_returns, tr_curves
|
||||
|
||||
|
||||
def create_trading_environment(year=2006):
|
||||
"""Construct a complete environment with reasonable defaults"""
|
||||
benchmark_returns, treasury_curves = load_market_data()
|
||||
@@ -51,10 +51,11 @@ def create_trading_environment(year=2006):
|
||||
period_end = end,
|
||||
capital_base = 100000.0
|
||||
)
|
||||
|
||||
|
||||
return trading_environment
|
||||
|
||||
def create_trade(sid, price, amount, datetime):
|
||||
row = zp.namedict({
|
||||
row = zp.ndict({
|
||||
'source_id' : "test_factory",
|
||||
'type' : zp.DATASOURCE_TYPE.TRADE,
|
||||
'sid' : sid,
|
||||
@@ -70,7 +71,7 @@ def get_next_trading_dt(current, interval, trading_calendar):
|
||||
next = next + interval
|
||||
if trading_calendar.is_market_hours(next):
|
||||
break
|
||||
|
||||
|
||||
return next
|
||||
|
||||
def create_trade_history(sid, prices, amounts, interval, trading_calendar):
|
||||
@@ -78,7 +79,7 @@ def create_trade_history(sid, prices, amounts, interval, trading_calendar):
|
||||
current = trading_calendar.first_open
|
||||
|
||||
for price, amount in zip(prices, amounts):
|
||||
|
||||
|
||||
trade = create_trade(sid, price, amount, current)
|
||||
trades.append(trade)
|
||||
current = get_next_trading_dt(current, interval, trading_calendar)
|
||||
@@ -87,11 +88,11 @@ def create_trade_history(sid, prices, amounts, interval, trading_calendar):
|
||||
return trades
|
||||
|
||||
def create_txn(sid, price, amount, datetime, btrid=None):
|
||||
txn = zp.namedict({
|
||||
'sid':sid,
|
||||
'amount':amount,
|
||||
'dt':datetime,
|
||||
'price':price,
|
||||
txn = zp.ndict({
|
||||
'sid' : sid,
|
||||
'amount' : amount,
|
||||
'dt' : datetime,
|
||||
'price' : price,
|
||||
})
|
||||
return txn
|
||||
|
||||
@@ -115,15 +116,15 @@ def create_returns(daycount, trading_calendar):
|
||||
test_range = []
|
||||
current = trading_calendar.first_open
|
||||
one_day = timedelta(days = 1)
|
||||
|
||||
for day in range(daycount):
|
||||
|
||||
for day in range(daycount):
|
||||
current = current + one_day
|
||||
if trading_calendar.is_trading_day(current):
|
||||
r = risk.DailyReturn(current, random.random())
|
||||
test_range.append(r)
|
||||
|
||||
|
||||
return test_range
|
||||
|
||||
|
||||
|
||||
def create_returns_from_range(trading_calendar):
|
||||
current = trading_calendar.first_open
|
||||
@@ -134,53 +135,53 @@ def create_returns_from_range(trading_calendar):
|
||||
r = risk.DailyReturn(current, random.random())
|
||||
test_range.append(r)
|
||||
current = get_next_trading_dt(current, one_day, trading_calendar)
|
||||
|
||||
|
||||
return test_range
|
||||
|
||||
|
||||
def create_returns_from_list(returns, trading_calendar):
|
||||
current = trading_calendar.first_open
|
||||
one_day = timedelta(days = 1)
|
||||
test_range = []
|
||||
|
||||
|
||||
#sometimes the range starts with a non-trading day.
|
||||
if not trading_calendar.is_trading_day(current):
|
||||
current = get_next_trading_dt(current, one_day, trading_calendar)
|
||||
|
||||
for return_val in returns:
|
||||
|
||||
for return_val in returns:
|
||||
r = risk.DailyReturn(current, return_val)
|
||||
test_range.append(r)
|
||||
current = get_next_trading_dt(current, one_day, trading_calendar)
|
||||
|
||||
|
||||
return test_range
|
||||
|
||||
def create_random_trade_source(sid, trade_count, trading_environment):
|
||||
# create the source
|
||||
source = RandomEquityTrades(sid, "rand-"+str(sid), trade_count)
|
||||
|
||||
|
||||
# make the period_end of trading_environment match
|
||||
cur = trading_environment.first_open
|
||||
one_day = timedelta(days = 1)
|
||||
for i in range(trade_count + 2):
|
||||
cur = get_next_trading_dt(cur, one_day, trading_environment)
|
||||
trading_environment.period_end = cur
|
||||
|
||||
|
||||
return source
|
||||
|
||||
|
||||
def create_daily_trade_source(sids, trade_count, trading_environment):
|
||||
|
||||
|
||||
"""
|
||||
creates trade_count trades for each sid in sids list.
|
||||
first trade will be on trading_environment.period_start, and daily
|
||||
thereafter for each sid. Thus, two sids should result in two trades per
|
||||
day.
|
||||
|
||||
creates trade_count trades for each sid in sids list.
|
||||
first trade will be on trading_environment.period_start, and daily
|
||||
thereafter for each sid. Thus, two sids should result in two trades per
|
||||
day.
|
||||
|
||||
Important side-effect: trading_environment.period_end will be modified
|
||||
to match the day of the final trade.
|
||||
to match the day of the final trade.
|
||||
"""
|
||||
return create_trade_source(
|
||||
sids,
|
||||
trade_count,
|
||||
timedelta(days=1),
|
||||
sids,
|
||||
trade_count,
|
||||
timedelta(days=1),
|
||||
trading_environment
|
||||
)
|
||||
|
||||
@@ -188,44 +189,45 @@ def create_daily_trade_source(sids, trade_count, trading_environment):
|
||||
def create_minutely_trade_source(sids, trade_count, trading_environment):
|
||||
|
||||
"""
|
||||
creates trade_count trades for each sid in sids list.
|
||||
first trade will be on trading_environment.period_start, and every minute
|
||||
thereafter for each sid. Thus, two sids should result in two trades per
|
||||
minute.
|
||||
creates trade_count trades for each sid in sids list.
|
||||
first trade will be on trading_environment.period_start, and every minute
|
||||
thereafter for each sid. Thus, two sids should result in two trades per
|
||||
minute.
|
||||
|
||||
Important side-effect: trading_environment.period_end will be modified
|
||||
to match the day of the final trade.
|
||||
to match the day of the final trade.
|
||||
"""
|
||||
return create_trade_source(
|
||||
sids,
|
||||
trade_count,
|
||||
timedelta(minutes=1),
|
||||
sids,
|
||||
trade_count,
|
||||
timedelta(minutes=1),
|
||||
trading_environment
|
||||
)
|
||||
|
||||
def create_trade_source(sids, trade_count, trade_time_increment, trading_environment):
|
||||
trade_history = []
|
||||
|
||||
price = [10.1] * trade_count
|
||||
volume = [100] * trade_count
|
||||
|
||||
for sid in sids:
|
||||
price = [10.1] * trade_count
|
||||
volume = [100] * trade_count
|
||||
start_date = trading_environment.first_open
|
||||
|
||||
generated_trades = create_trade_history(
|
||||
sid,
|
||||
price,
|
||||
volume,
|
||||
trade_time_increment,
|
||||
trading_environment
|
||||
generated_trades = create_trade_history(
|
||||
sid,
|
||||
price,
|
||||
volume,
|
||||
trade_time_increment,
|
||||
trading_environment
|
||||
)
|
||||
|
||||
|
||||
trade_history.extend(generated_trades)
|
||||
|
||||
|
||||
trade_history = sorted(trade_history, key=lambda(x): x.dt)
|
||||
|
||||
|
||||
#set the trading environment's end to same dt as the last trade in the
|
||||
#history.
|
||||
trading_environment.period_end = trade_history[-1].dt
|
||||
|
||||
|
||||
source = SpecificEquityTrades("flat", trade_history)
|
||||
return source
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
"""
|
||||
Small classes to assist with timezone calculations, LOGGER configuration,
|
||||
and other common operations.
|
||||
"""
|
||||
|
||||
import logging
|
||||
import logging.config
|
||||
|
||||
def configure_logging():
|
||||
logging.config.fileConfig(
|
||||
'logging.cfg',
|
||||
disable_existing_loggers = False
|
||||
)
|
||||
@@ -1,28 +0,0 @@
|
||||
"""
|
||||
Small classes to assist with timezone calculations, LOGGER configuration,
|
||||
and other common operations.
|
||||
"""
|
||||
|
||||
import datetime
|
||||
import pytz
|
||||
import logging
|
||||
import logging.handlers
|
||||
|
||||
LOGGER = logging.getLogger('ZiplineLogger')
|
||||
|
||||
def configure_logging(loglevel=logging.DEBUG):
|
||||
"""
|
||||
Configures zipline.util.LOGGER to write a rotating file
|
||||
(10M per file, 5 files) to `` /var/log/zipline.log ``.
|
||||
"""
|
||||
LOGGER.setLevel(loglevel)
|
||||
handler = logging.handlers.RotatingFileHandler(
|
||||
"/var/log/zipline/{lfn}.log".format(lfn="zipline"),
|
||||
maxBytes=10*1024*1024, backupCount=5
|
||||
)
|
||||
handler.setFormatter(logging.Formatter(
|
||||
"%(asctime)s %(levelname)s %(filename)s %(funcName)s - %(message)s",
|
||||
"%Y-%m-%d %H:%M:%S %Z")
|
||||
)
|
||||
LOGGER.addHandler(handler)
|
||||
LOGGER.info("logging started...")
|
||||
@@ -31,79 +31,6 @@ def FrameExceptionFactory(name):
|
||||
|
||||
return InvalidFrame
|
||||
|
||||
class namedict(MutableMapping):
|
||||
"""
|
||||
|
||||
Namedicts are dict like objects that have fields accessible by attribute lookup
|
||||
as well as being indexable and iterable::
|
||||
|
||||
HEARTBEAT_PROTOCOL = namedict({
|
||||
'REQ' : b'\x01',
|
||||
'REP' : b'\x02',
|
||||
})
|
||||
|
||||
HEARTBEAT_PROTOCOL.REQ # syntactic sugar
|
||||
HEARTBEAT_PROTOCOL.REP # oh suga suga
|
||||
|
||||
For more complex structs use collections.namedtuple:
|
||||
"""
|
||||
|
||||
def __init__(self, dct=None):
|
||||
if(dct):
|
||||
self.__dict__.update(dct)
|
||||
|
||||
def __setitem__(self, key, value):
|
||||
"""
|
||||
Required for use by pymongo as_class parameter to find.
|
||||
"""
|
||||
if(key == '_id'):
|
||||
self.__dict__['id'] = value
|
||||
else:
|
||||
self.__dict__[key] = value
|
||||
|
||||
def __getitem__(self, key):
|
||||
return self.__dict__[key]
|
||||
|
||||
def __delitem__(self, key):
|
||||
del self.__dict__[key]
|
||||
|
||||
def __iter__(self):
|
||||
return self.__dict__.iterkeys()
|
||||
|
||||
def __len__(self):
|
||||
return len(self.__dict__)
|
||||
|
||||
def keys(self):
|
||||
return self.__dict__.keys()
|
||||
|
||||
def as_dict(self):
|
||||
# shallow copy is O(n)
|
||||
return copy.copy(self.__dict__)
|
||||
|
||||
def delete(self, key):
|
||||
del(self.__dict__[key])
|
||||
|
||||
def merge(self, other_nd):
|
||||
assert isinstance(other_nd, namedict)
|
||||
self.__dict__.update(other_nd.__dict__)
|
||||
|
||||
def __repr__(self):
|
||||
return "namedict: " + str(self.__dict__)
|
||||
|
||||
def __eq__(self, other):
|
||||
# !!!!!!!!!!!!!!!!!!!!
|
||||
# !!!! DANGEROUS !!!!!
|
||||
# !!!!!!!!!!!!!!!!!!!!
|
||||
return other != None and self.__dict__ == other.__dict__
|
||||
|
||||
def has_attr(self, name):
|
||||
return self.__dict__.has_key(name)
|
||||
|
||||
def as_series(self):
|
||||
s = pandas.Series(self.__dict__)
|
||||
s.name = self.sid
|
||||
return s
|
||||
|
||||
class ndict(MutableMapping):
|
||||
"""
|
||||
Xtreme Namedicts 2.0
|
||||
@@ -123,6 +50,13 @@ class ndict(MutableMapping):
|
||||
# Abstact Overloads
|
||||
# -----------------
|
||||
|
||||
def __setattr__(self, key, value):
|
||||
if 'ndict' in key or key == 'cls':
|
||||
MutableMapping.__setattr__(self, key, value)
|
||||
else:
|
||||
self.__internal[key] = value
|
||||
return value
|
||||
|
||||
def __setitem__(self, key, value):
|
||||
"""
|
||||
Required for use by pymongo as_class parameter to find.
|
||||
@@ -132,7 +66,6 @@ class ndict(MutableMapping):
|
||||
else:
|
||||
self.__internal[key] = value
|
||||
|
||||
|
||||
def __getattr__(self, key):
|
||||
if key in self.cls:
|
||||
return self.__dict__[key]
|
||||
@@ -219,3 +152,23 @@ class ndict(MutableMapping):
|
||||
#return False
|
||||
|
||||
#return True
|
||||
|
||||
# This is not neccesarily the most intuitive construction, but
|
||||
# we're aiming for raw performance rather than readability. So
|
||||
# we do things that we would not normally do in business logic.
|
||||
def namelookup(dct):
|
||||
ks = dct.keys()
|
||||
vs = dct.values()
|
||||
dct = {}
|
||||
class _lookup:
|
||||
__slots__ = ks
|
||||
def __init__(self):
|
||||
for k, v in zip(ks, vs):
|
||||
setattr(self,k,v)
|
||||
self.__setattr__ = self.locked
|
||||
def locked(self,k,v):
|
||||
raise Exception('Name lookups are fixed at init.')
|
||||
def __repr__(self):
|
||||
return '<namelookup %s>' % self.__slots__
|
||||
del dct
|
||||
return _lookup()
|
||||
|
||||
Reference in New Issue
Block a user