From 4b33d6ea4b0de94c5c093594a81983be1806cff4 Mon Sep 17 00:00:00 2001 From: Eddie Hebert Date: Thu, 25 Apr 2013 12:50:40 -0400 Subject: [PATCH] TST: Ensure that create_trade_history uses midnight for daily trades. Prepare for implementation of backtest loop that depends on daily trades being grouped by midnight. --- tests/test_perf_tracking.py | 8 +++++--- zipline/utils/factory.py | 8 +++++++- 2 files changed, 12 insertions(+), 4 deletions(-) diff --git a/tests/test_perf_tracking.py b/tests/test_perf_tracking.py index c7209de4..a99ecd90 100644 --- a/tests/test_perf_tracking.py +++ b/tests/test_perf_tracking.py @@ -80,6 +80,7 @@ class TestDividendPerformance(unittest.TestCase): ) self.assertEqual(after.hour, 13) + @trading.use_environment(trading.TradingEnvironment()) def test_long_position_receives_dividend(self): #post some trades in the market events = factory.create_trade_history( @@ -126,11 +127,12 @@ class TestDividendPerformance(unittest.TestCase): perf_messages, risk = perf_tracker.handle_simulation_end() results.append(perf_messages[0]) - self.assertEqual(results[0]['daily_perf']['period_open'], events[0].dt) + self.assertEqual( + results[0]['daily_perf']['period_open'], + trading.environment.get_open_and_close(events[0].dt)[0]) self.assertEqual( results[-1]['daily_perf']['period_open'], - events[-1].dt - ) + trading.environment.get_open_and_close(events[-1].dt)[0]) self.assertEqual(len(results), 5) cumulative_returns = \ diff --git a/zipline/utils/factory.py b/zipline/utils/factory.py index 79c04333..23f4d22e 100644 --- a/zipline/utils/factory.py +++ b/zipline/utils/factory.py @@ -145,8 +145,14 @@ def create_trade_history(sid, prices, amounts, interval, sim_params, trades = [] current = sim_params.first_open + oneday = timedelta(days=1) + use_midnight = interval >= oneday for price, amount in zip(prices, amounts): - trade = create_trade(sid, price, amount, current, source_id) + if use_midnight: + trade_dt = current.replace(hour=0, minute=0) + else: + trade_dt = current + trade = create_trade(sid, price, amount, trade_dt, source_id) trades.append(trade) current = get_next_trading_dt(current, interval)