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:
jfkirk
2015-09-10 11:53:28 -04:00
parent 661314ce49
commit dc964a7e7d
45 changed files with 1484 additions and 1173 deletions
+36 -21
View File
@@ -19,10 +19,12 @@ import random
from six.moves import range, map
from nose_parameterized import parameterized
from unittest import TestCase
from functools import partial
from collections import namedtuple
import numpy as np
from zipline.finance.trading import TradingEnvironment, with_environment
from zipline.finance.trading import TradingEnvironment
import zipline.utils.events
from zipline.utils.events import (
EventRule,
@@ -161,7 +163,7 @@ class TestEventManager(TestCase):
class CountingRule(Always):
count = 0
def should_trigger(self, dt):
def should_trigger(self, dt, env):
CountingRule.count += 1
return True
@@ -170,7 +172,10 @@ class TestEventManager(TestCase):
Event(r(), lambda context, data: None)
)
self.em.handle_data(None, None, datetime.datetime.now())
mock_algo_class = namedtuple('FakeAlgo', ['trading_environment'])
mock_algo = mock_algo_class(trading_environment="fake_env")
self.em.handle_data(mock_algo, None, datetime.datetime.now(),
mock_algo.trading_environment)
self.assertEqual(CountingRule.count, 5)
@@ -182,11 +187,10 @@ class TestEventRule(TestCase):
def test_not_implemented(self):
with self.assertRaises(NotImplementedError):
super(Always, Always()).should_trigger('a')
super(Always, Always()).should_trigger('a', env=None)
@with_environment()
def minutes_for_days(env=None):
def minutes_for_days():
"""
500 randomly selected days.
This is used to make sure our test coverage is unbaised towards any rules.
@@ -202,6 +206,7 @@ def minutes_for_days(env=None):
Iterating over this yeilds a single day, iterating over the day yields
the minutes for that day.
"""
env = TradingEnvironment()
random.seed('deterministic')
return ((env.market_minutes_for_day(random.choice(env.trading_days)),)
for _ in range(500))
@@ -210,7 +215,7 @@ def minutes_for_days(env=None):
class RuleTestCase(TestCase):
@classmethod
def setUpClass(cls):
cls.env = TradingEnvironment.instance()
cls.env = TradingEnvironment()
cls.class_ = None # Mark that this is the base class.
def test_completeness(self):
@@ -256,17 +261,18 @@ class TestStatelessRules(RuleTestCase):
@parameterized.expand(minutes_for_days())
def test_Always(self, ms):
should_trigger = Always().should_trigger
self.assertTrue(all(map(should_trigger, ms)))
should_trigger = partial(Always().should_trigger, env=self.env)
self.assertTrue(all(map(partial(should_trigger, env=self.env), ms)))
@parameterized.expand(minutes_for_days())
def test_Never(self, ms):
should_trigger = Never().should_trigger
should_trigger = partial(Never().should_trigger, env=self.env)
self.assertFalse(any(map(should_trigger, ms)))
@parameterized.expand(minutes_for_days())
def test_AfterOpen(self, ms):
should_trigger = AfterOpen(minutes=5, hours=1).should_trigger
should_trigger = partial(AfterOpen(minutes=5, hours=1).should_trigger,
env=self.env)
for m in islice(ms, 64):
# Check the first 64 minutes of data.
# We use 64 because the offset is from market open
@@ -280,20 +286,23 @@ class TestStatelessRules(RuleTestCase):
@parameterized.expand(minutes_for_days())
def test_BeforeClose(self, ms):
ms = list(ms)
should_trigger = BeforeClose(hours=1, minutes=5).should_trigger
should_trigger = partial(
BeforeClose(hours=1, minutes=5).should_trigger, env=self.env
)
for m in ms[0:-66]:
self.assertFalse(should_trigger(m))
for m in ms[-66:]:
self.assertTrue(should_trigger(m))
def test_NotHalfDay(self):
should_trigger = NotHalfDay().should_trigger
should_trigger = partial(NotHalfDay().should_trigger, env=self.env)
self.assertTrue(should_trigger(FULL_DAY))
self.assertFalse(should_trigger(HALF_DAY))
@parameterized.expand(param_range(MAX_WEEK_RANGE))
def test_NthTradingDayOfWeek(self, n):
should_trigger = NthTradingDayOfWeek(n).should_trigger
should_trigger = partial(NthTradingDayOfWeek(n).should_trigger,
env=self.env)
prev_day = self.sept_week[0].date()
n_tdays = 0
for m in self.sept_week:
@@ -308,7 +317,9 @@ class TestStatelessRules(RuleTestCase):
@parameterized.expand(param_range(MAX_WEEK_RANGE))
def test_NDaysBeforeLastTradingDayOfWeek(self, n):
should_trigger = NDaysBeforeLastTradingDayOfWeek(n).should_trigger
should_trigger = partial(
NDaysBeforeLastTradingDayOfWeek(n).should_trigger, env=self.env
)
for m in self.sept_week:
if should_trigger(m):
n_tdays = 0
@@ -323,7 +334,8 @@ class TestStatelessRules(RuleTestCase):
@parameterized.expand(param_range(MAX_MONTH_RANGE))
def test_NthTradingDayOfMonth(self, n):
should_trigger = NthTradingDayOfMonth(n).should_trigger
should_trigger = partial(NthTradingDayOfMonth(n).should_trigger,
env=self.env)
for n_tdays, d in enumerate(self.sept_days):
for m in self.env.market_minutes_for_day(d):
if should_trigger(m):
@@ -333,7 +345,9 @@ class TestStatelessRules(RuleTestCase):
@parameterized.expand(param_range(MAX_MONTH_RANGE))
def test_NDaysBeforeLastTradingDayOfMonth(self, n):
should_trigger = NDaysBeforeLastTradingDayOfMonth(n).should_trigger
should_trigger = partial(
NDaysBeforeLastTradingDayOfMonth(n).should_trigger, env=self.env
)
for n_days_before, d in enumerate(reversed(self.sept_days)):
for m in self.env.market_minutes_for_day(d):
if should_trigger(m):
@@ -347,10 +361,11 @@ class TestStatelessRules(RuleTestCase):
rule2 = Never()
composed = rule1 & rule2
should_trigger = partial(composed.should_trigger, env=self.env)
self.assertIsInstance(composed, ComposedRule)
self.assertIs(composed.first, rule1)
self.assertIs(composed.second, rule2)
self.assertFalse(any(map(composed.should_trigger, ms)))
self.assertFalse(any(map(should_trigger, ms)))
class TestStatefulRules(RuleTestCase):
@@ -369,14 +384,14 @@ class TestStatefulRules(RuleTestCase):
"""
count = 0
def should_trigger(self, dt):
st = self.rule.should_trigger(dt)
def should_trigger(self, dt, env):
st = self.rule.should_trigger(dt, env)
if st:
self.count += 1
return st
rule = RuleCounter(OncePerDay())
for m in ms:
rule.should_trigger(m)
rule.should_trigger(m, env=self.env)
self.assertEqual(rule.count, 1)