mirror of
https://github.com/wassname/catalyst.git
synced 2026-09-09 11:19:23 +08:00
MAINT: Removes the ability to reference a global TradingEnvironment
This commit removes the ability to reference a shared TradingEnvironment through the zipline.finance.trading module. In place, the classes that require a TradingEnvironment, or its child AssetFinder, contain their own references to those objects. This commit also adds serialization utilities that allow for the pickling/unpickling of objects without unintentionally their TradingEnvironments or AssetFinders.
This commit is contained in:
+82
-64
@@ -6,8 +6,7 @@ from unittest import TestCase
|
||||
from zipline.algorithm import TradingAlgorithm
|
||||
from zipline.errors import TradingControlViolation
|
||||
from zipline.sources import SpecificEquityTrades
|
||||
from zipline.finance import trading
|
||||
from zipline.finance.trading import with_environment
|
||||
from zipline.finance.trading import TradingEnvironment
|
||||
from zipline.utils.test_utils import (
|
||||
setup_logger, teardown_logger, security_list_copy, add_security_data,)
|
||||
from zipline.utils import factory
|
||||
@@ -19,7 +18,7 @@ LEVERAGED_ETFS = load_from_directory('leveraged_etf_list')
|
||||
|
||||
class RestrictedAlgoWithCheck(TradingAlgorithm):
|
||||
def initialize(self, symbol):
|
||||
self.rl = SecurityListSet(self.get_datetime)
|
||||
self.rl = SecurityListSet(self.get_datetime, self.asset_finder)
|
||||
self.set_do_not_order_list(self.rl.leveraged_etf_list)
|
||||
self.order_count = 0
|
||||
self.sid = self.symbol(symbol)
|
||||
@@ -34,7 +33,7 @@ class RestrictedAlgoWithCheck(TradingAlgorithm):
|
||||
|
||||
class RestrictedAlgoWithoutCheck(TradingAlgorithm):
|
||||
def initialize(self, symbol):
|
||||
self.rl = SecurityListSet(self.get_datetime)
|
||||
self.rl = SecurityListSet(self.get_datetime, self.asset_finder)
|
||||
self.set_do_not_order_list(self.rl.leveraged_etf_list)
|
||||
self.order_count = 0
|
||||
self.sid = self.symbol(symbol)
|
||||
@@ -46,7 +45,7 @@ class RestrictedAlgoWithoutCheck(TradingAlgorithm):
|
||||
|
||||
class IterateRLAlgo(TradingAlgorithm):
|
||||
def initialize(self, symbol):
|
||||
self.rl = SecurityListSet(self.get_datetime)
|
||||
self.rl = SecurityListSet(self.get_datetime, self.asset_finder)
|
||||
self.set_do_not_order_list(self.rl.leveraged_etf_list)
|
||||
self.order_count = 0
|
||||
self.sid = self.symbol(symbol)
|
||||
@@ -60,6 +59,12 @@ class IterateRLAlgo(TradingAlgorithm):
|
||||
|
||||
class SecurityListTestCase(TestCase):
|
||||
|
||||
@classmethod
|
||||
def setUpClass(cls):
|
||||
cls.env = TradingEnvironment()
|
||||
cls.env.write_data(equities_identifiers=['AAPL', 'GOOG', 'BZQ',
|
||||
'URTY', 'JFT'])
|
||||
|
||||
def setUp(self, env=None):
|
||||
|
||||
self.extra_knowledge_date = \
|
||||
@@ -69,43 +74,38 @@ class SecurityListTestCase(TestCase):
|
||||
|
||||
setup_logger(self)
|
||||
|
||||
trading.environment = trading.TradingEnvironment()
|
||||
|
||||
def tearDown(self):
|
||||
teardown_logger(self)
|
||||
|
||||
def test_iterate_over_rl(self):
|
||||
sim_params = factory.create_simulation_parameters(
|
||||
start=list(LEVERAGED_ETFS.keys())[0], num_days=4)
|
||||
trading.environment.write_data(equities_identifiers=['BZQ'])
|
||||
start=list(LEVERAGED_ETFS.keys())[0], num_days=4, env=self.env)
|
||||
trade_history = factory.create_trade_history(
|
||||
'BZQ',
|
||||
[10.0, 10.0, 11.0, 11.0],
|
||||
[100, 100, 100, 300],
|
||||
timedelta(days=1),
|
||||
sim_params
|
||||
sim_params,
|
||||
env=self.env
|
||||
)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history)
|
||||
algo = IterateRLAlgo(symbol='BZQ', sim_params=sim_params)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history,
|
||||
env=self.env)
|
||||
algo = IterateRLAlgo(symbol='BZQ', sim_params=sim_params, env=self.env)
|
||||
algo.run(self.source)
|
||||
self.assertTrue(algo.found)
|
||||
|
||||
@with_environment()
|
||||
def test_security_list(self, env=None):
|
||||
def test_security_list(self):
|
||||
|
||||
# set the knowledge date to the first day of the
|
||||
# leveraged etf knowledge date.
|
||||
def get_datetime():
|
||||
return list(LEVERAGED_ETFS.keys())[0]
|
||||
|
||||
env.write_data(equities_identifiers=['AAPL', 'GOOG', 'BZQ',
|
||||
'URTY', 'JFT'])
|
||||
|
||||
rl = SecurityListSet(get_datetime)
|
||||
rl = SecurityListSet(get_datetime, self.env.asset_finder)
|
||||
# assert that a sample from the leveraged list are in restricted
|
||||
should_exist = [
|
||||
asset.sid for asset in
|
||||
[env.asset_finder.lookup_symbol(
|
||||
[self.env.asset_finder.lookup_symbol(
|
||||
symbol,
|
||||
as_of_date=self.extra_knowledge_date)
|
||||
for symbol in ["BZQ", "URTY", "JFT"]]
|
||||
@@ -116,7 +116,7 @@ class SecurityListTestCase(TestCase):
|
||||
# assert that a sample of allowed stocks are not in restricted
|
||||
shouldnt_exist = [
|
||||
asset.sid for asset in
|
||||
[env.asset_finder.lookup_symbol(
|
||||
[self.env.asset_finder.lookup_symbol(
|
||||
symbol,
|
||||
as_of_date=self.extra_knowledge_date)
|
||||
for symbol in ["AAPL", "GOOG"]]
|
||||
@@ -124,18 +124,15 @@ class SecurityListTestCase(TestCase):
|
||||
for sid in shouldnt_exist:
|
||||
self.assertNotIn(sid, rl.leveraged_etf_list)
|
||||
|
||||
@with_environment()
|
||||
def test_security_add(self, env=None):
|
||||
def test_security_add(self):
|
||||
def get_datetime():
|
||||
return datetime(2015, 1, 27, tzinfo=pytz.utc)
|
||||
with security_list_copy():
|
||||
add_security_data(['AAPL', 'GOOG'], [])
|
||||
env.write_data(equities_identifiers=['AAPL', 'GOOG',
|
||||
'BZQ', 'URTY'])
|
||||
rl = SecurityListSet(get_datetime)
|
||||
rl = SecurityListSet(get_datetime, self.env.asset_finder)
|
||||
should_exist = [
|
||||
asset.sid for asset in
|
||||
[env.asset_finder.lookup_symbol(
|
||||
[self.env.asset_finder.lookup_symbol(
|
||||
symbol,
|
||||
as_of_date=self.extra_knowledge_date
|
||||
) for symbol in ["AAPL", "GOOG", "BZQ", "URTY"]]
|
||||
@@ -147,57 +144,67 @@ class SecurityListTestCase(TestCase):
|
||||
with security_list_copy():
|
||||
def get_datetime():
|
||||
return datetime(2015, 1, 27, tzinfo=pytz.utc)
|
||||
trading.environment.write_data(equities_identifiers=['BZQ',
|
||||
'URTY'])
|
||||
rl = SecurityListSet(get_datetime)
|
||||
rl = SecurityListSet(get_datetime, self.env.asset_finder)
|
||||
self.assertNotIn("BZQ", rl.leveraged_etf_list)
|
||||
self.assertNotIn("URTY", rl.leveraged_etf_list)
|
||||
|
||||
def test_algo_without_rl_violation_via_check(self):
|
||||
sim_params = factory.create_simulation_parameters(
|
||||
start=list(LEVERAGED_ETFS.keys())[0], num_days=4)
|
||||
trading.environment.write_data(equities_identifiers=['BZQ'])
|
||||
start=list(LEVERAGED_ETFS.keys())[0], num_days=4,
|
||||
env=self.env)
|
||||
trade_history = factory.create_trade_history(
|
||||
'BZQ',
|
||||
[10.0, 10.0, 11.0, 11.0],
|
||||
[100, 100, 100, 300],
|
||||
timedelta(days=1),
|
||||
sim_params
|
||||
sim_params,
|
||||
env=self.env
|
||||
)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history,
|
||||
env=self.env)
|
||||
|
||||
algo = RestrictedAlgoWithCheck(symbol='BZQ', sim_params=sim_params)
|
||||
algo = RestrictedAlgoWithCheck(symbol='BZQ',
|
||||
sim_params=sim_params,
|
||||
env=self.env)
|
||||
algo.run(self.source)
|
||||
|
||||
def test_algo_without_rl_violation(self):
|
||||
sim_params = factory.create_simulation_parameters(
|
||||
start=list(LEVERAGED_ETFS.keys())[0], num_days=4)
|
||||
trading.environment.write_data(equities_identifiers=['AAPL'])
|
||||
start=list(LEVERAGED_ETFS.keys())[0], num_days=4,
|
||||
env=self.env)
|
||||
trade_history = factory.create_trade_history(
|
||||
'AAPL',
|
||||
[10.0, 10.0, 11.0, 11.0],
|
||||
[100, 100, 100, 300],
|
||||
timedelta(days=1),
|
||||
sim_params
|
||||
sim_params,
|
||||
env=self.env
|
||||
)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history)
|
||||
algo = RestrictedAlgoWithoutCheck(symbol='AAPL', sim_params=sim_params)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history,
|
||||
env=self.env)
|
||||
algo = RestrictedAlgoWithoutCheck(symbol='AAPL',
|
||||
sim_params=sim_params,
|
||||
env=self.env)
|
||||
algo.run(self.source)
|
||||
|
||||
def test_algo_with_rl_violation(self):
|
||||
sim_params = factory.create_simulation_parameters(
|
||||
start=list(LEVERAGED_ETFS.keys())[0], num_days=4)
|
||||
trading.environment.write_data(equities_identifiers=['BZQ', 'JFT'])
|
||||
start=list(LEVERAGED_ETFS.keys())[0], num_days=4,
|
||||
env=self.env)
|
||||
trade_history = factory.create_trade_history(
|
||||
'BZQ',
|
||||
[10.0, 10.0, 11.0, 11.0],
|
||||
[100, 100, 100, 300],
|
||||
timedelta(days=1),
|
||||
sim_params
|
||||
sim_params,
|
||||
env=self.env
|
||||
)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history,
|
||||
env=self.env)
|
||||
|
||||
algo = RestrictedAlgoWithoutCheck(symbol='BZQ', sim_params=sim_params)
|
||||
algo = RestrictedAlgoWithoutCheck(symbol='BZQ',
|
||||
sim_params=sim_params,
|
||||
env=self.env)
|
||||
with self.assertRaises(TradingControlViolation) as ctx:
|
||||
algo.run(self.source)
|
||||
|
||||
@@ -209,11 +216,15 @@ class SecurityListTestCase(TestCase):
|
||||
[10.0, 10.0, 11.0, 11.0],
|
||||
[100, 100, 100, 300],
|
||||
timedelta(days=1),
|
||||
sim_params
|
||||
sim_params,
|
||||
env=self.env
|
||||
)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history,
|
||||
env=self.env)
|
||||
|
||||
algo = RestrictedAlgoWithoutCheck(symbol='JFT', sim_params=sim_params)
|
||||
algo = RestrictedAlgoWithoutCheck(symbol='JFT',
|
||||
sim_params=sim_params,
|
||||
env=self.env)
|
||||
with self.assertRaises(TradingControlViolation) as ctx:
|
||||
algo.run(self.source)
|
||||
|
||||
@@ -222,17 +233,21 @@ class SecurityListTestCase(TestCase):
|
||||
def test_algo_with_rl_violation_after_knowledge_date(self):
|
||||
sim_params = factory.create_simulation_parameters(
|
||||
start=list(
|
||||
LEVERAGED_ETFS.keys())[0] + timedelta(days=7), num_days=5)
|
||||
trading.environment.write_data(equities_identifiers=['BZQ'])
|
||||
LEVERAGED_ETFS.keys())[0] + timedelta(days=7), num_days=5,
|
||||
env=self.env)
|
||||
trade_history = factory.create_trade_history(
|
||||
'BZQ',
|
||||
[10.0, 10.0, 11.0, 11.0],
|
||||
[100, 100, 100, 300],
|
||||
timedelta(days=1),
|
||||
sim_params
|
||||
sim_params,
|
||||
env=self.env
|
||||
)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history)
|
||||
algo = RestrictedAlgoWithoutCheck(symbol='BZQ', sim_params=sim_params)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history,
|
||||
env=self.env)
|
||||
algo = RestrictedAlgoWithoutCheck(symbol='BZQ',
|
||||
sim_params=sim_params,
|
||||
env=self.env)
|
||||
with self.assertRaises(TradingControlViolation) as ctx:
|
||||
algo.run(self.source)
|
||||
|
||||
@@ -255,12 +270,13 @@ class SecurityListTestCase(TestCase):
|
||||
[10.0, 10.0, 11.0, 11.0],
|
||||
[100, 100, 100, 300],
|
||||
timedelta(days=1),
|
||||
sim_params
|
||||
sim_params,
|
||||
env=self.env,
|
||||
)
|
||||
trading.environment.write_data(equities_identifiers=['BZQ'])
|
||||
self.source = SpecificEquityTrades(event_list=trade_history)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history,
|
||||
env=self.env)
|
||||
algo = RestrictedAlgoWithoutCheck(
|
||||
symbol='BZQ', sim_params=sim_params)
|
||||
symbol='BZQ', sim_params=sim_params, env=self.env)
|
||||
with self.assertRaises(TradingControlViolation) as ctx:
|
||||
algo.run(self.source)
|
||||
|
||||
@@ -273,18 +289,19 @@ class SecurityListTestCase(TestCase):
|
||||
add_security_data([], ['BZQ'])
|
||||
sim_params = factory.create_simulation_parameters(
|
||||
start=self.extra_knowledge_date, num_days=3)
|
||||
trading.environment.write_data(equities_identifiers=['BZQ'])
|
||||
|
||||
trade_history = factory.create_trade_history(
|
||||
'BZQ',
|
||||
[10.0, 10.0, 11.0, 11.0],
|
||||
[100, 100, 100, 300],
|
||||
timedelta(days=1),
|
||||
sim_params
|
||||
sim_params,
|
||||
env=self.env,
|
||||
)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history,
|
||||
env=self.env)
|
||||
algo = RestrictedAlgoWithoutCheck(
|
||||
symbol='BZQ', sim_params=sim_params
|
||||
symbol='BZQ', sim_params=sim_params, env=self.env
|
||||
)
|
||||
algo.run(self.source)
|
||||
|
||||
@@ -293,17 +310,18 @@ class SecurityListTestCase(TestCase):
|
||||
add_security_data(['AAPL'], [])
|
||||
sim_params = factory.create_simulation_parameters(
|
||||
start=self.trading_day_before_first_kd, num_days=4)
|
||||
trading.environment.write_data(equities_identifiers=['AAPL'])
|
||||
trade_history = factory.create_trade_history(
|
||||
'AAPL',
|
||||
[10.0, 10.0, 11.0, 11.0],
|
||||
[100, 100, 100, 300],
|
||||
timedelta(days=1),
|
||||
sim_params
|
||||
sim_params,
|
||||
env=self.env
|
||||
)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history)
|
||||
self.source = SpecificEquityTrades(event_list=trade_history,
|
||||
env=self.env)
|
||||
algo = RestrictedAlgoWithoutCheck(
|
||||
symbol='AAPL', sim_params=sim_params)
|
||||
symbol='AAPL', sim_params=sim_params, env=self.env)
|
||||
with self.assertRaises(TradingControlViolation) as ctx:
|
||||
algo.run(self.source)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user