diff --git a/catalyst/exchange/data_portal_exchange.py b/catalyst/exchange/data_portal_exchange.py index 3dac7678..c0af0e32 100644 --- a/catalyst/exchange/data_portal_exchange.py +++ b/catalyst/exchange/data_portal_exchange.py @@ -29,7 +29,7 @@ from catalyst.exchange.exchange_errors import ( ExchangeRequestError, ExchangeBarDataError, BundleNotFoundError, PricingDataBeforeTradingError, - PricingDataNotLoadedError) + PricingDataNotLoadedError, InvalidHistoryFrequencyError) from catalyst.utils.paths import data_path log = Logger('DataPortalExchange') @@ -154,8 +154,9 @@ class DataPortalExchangeBase(DataPortal): try: if isinstance(assets, TradingPair): exchange = self.exchanges[assets.exchange] - return self.get_exchange_spot_value( - exchange, assets, field, dt, data_frequency) + spot_values = self.get_exchange_spot_value( + exchange, [assets], field, dt, data_frequency) + return spot_values[0] else: exchange_assets = dict() @@ -306,6 +307,15 @@ class DataPortalExchangeBacktest(DataPortalExchangeBase): @staticmethod def find_most_recent_time(bundle_name): + """ + Find most recent "time folder" for a given bundle. + + :param bundle_name: + The name of the targeted bundle. + + :return folder: + The name of the time folder. + """ try: bundle_folders = os.listdir( data_path([bundle_name]), @@ -326,6 +336,33 @@ class DataPortalExchangeBacktest(DataPortalExchangeBase): else: return None + def _get_reader(self, data_frequency, exchange_name): + """ + Pick from a collection of readers based of exchange name and frequency. + + :param data_frequency: + The reader frequency: minute, 5-minute, daily. + + :param exchange_name: + The exchange name. + + :return reader: + A reader object. + """ + if data_frequency == 'minute': + reader = self.minute_bar_readers[exchange_name] + elif data_frequency == '5-minute': + reader = self.five_minute_bar_readers[exchange_name] + elif data_frequency == 'daily': + reader = self.daily_bar_readers[exchange_name] + else: + raise InvalidHistoryFrequencyError(frequency=data_frequency) + + if reader is None: + raise ValueError('reader not found') + + return reader + def get_exchange_history_window(self, exchange, assets, @@ -335,33 +372,31 @@ class DataPortalExchangeBacktest(DataPortalExchangeBase): field, data_frequency, ffill=True): - if data_frequency == 'minute' or data_frequency == '1m': - reader = self.minute_bar_readers[exchange.name] + reader = self._get_reader(data_frequency, exchange.name) + if data_frequency == 'minute': dts = self.trading_calendar.minutes_window( end_dt, -bar_count ) self.ensure_after_first_day(dts[0], assets) - elif data_frequency == '5-minute' or data_frequency == '5m': - reader = self.five_minute_bar_readers[exchange.name] - elif data_frequency == 'daily' or data_frequency == '1d': - reader = self.daily_bar_readers[exchange.name] - + elif data_frequency == 'daily': session = self.trading_calendar.minute_to_session_label(end_dt) dts = self._get_days_for_window(session, bar_count) self.ensure_after_first_day(dts[0], assets) + else: - raise ValueError('Unsupported frequency') + raise InvalidHistoryFrequencyError(frequency=data_frequency) try: values = reader.load_raw_arrays( - [field], - dts[0], - dts[-1], - assets, + fields=[field], + start_dt=dts[0], + end_dt=dts[-1], + sids=[asset.sid for asset in assets], )[0] + except Exception: raise PricingDataNotLoadedError( field=field, @@ -393,50 +428,25 @@ class DataPortalExchangeBacktest(DataPortalExchangeBase): def get_exchange_spot_value(self, exchange, assets, field, dt, data_frequency): - if data_frequency == 'minute' or data_frequency == '1m': - reader = self.minute_bar_readers[exchange.name] - elif data_frequency == '5-minute' or data_frequency == '5m': - reader = self.five_minute_bar_readers[exchange.name] - elif data_frequency == 'daily' or data_frequency == '1d': - reader = self.daily_bar_readers[exchange.name] - else: - raise ValueError('Unsupported frequency') + reader = self._get_reader(data_frequency, exchange.name) - if isinstance(assets, TradingPair): - self.ensure_after_first_day(dt, [assets]) + self.ensure_after_first_day(dt, assets) + values = [] + for asset in assets: try: value = reader.get_value( - sid=assets.sid, + sid=asset.sid, dt=dt, field=field ) - return value + values.append(value) except Exception: raise PricingDataNotLoadedError( field=field, - first_trading_day=self._get_first_trading_day([assets]), + first_trading_day=self._get_first_trading_day(assets), exchange=exchange.name, - symbols=assets.symbol, + symbols=[asset.symbol for asset in assets], ) - else: - self.ensure_after_first_day(dt, assets) - values = [] - for asset in assets: - try: - value = reader.get_value( - sid=asset.sid, - dt=dt, - field=field - ) - values.append(value) - except Exception: - raise PricingDataNotLoadedError( - field=field, - first_trading_day=self._get_first_trading_day(assets), - exchange=exchange.name, - symbols=[asset.symbol for asset in assets], - ) - - return values + return values diff --git a/catalyst/exchange/exchange.py b/catalyst/exchange/exchange.py index 361413a5..76319b04 100644 --- a/catalyst/exchange/exchange.py +++ b/catalyst/exchange/exchange.py @@ -306,19 +306,14 @@ class Exchange: '1D', '7D', '14D', '1M' """ if field not in BASE_FIELDS: - raise KeyError('Invalid column: ' + str(field)) + raise KeyError('Invalid column: {}'.format(field)) - if isinstance(assets, collections.Iterable): - values = list() - for asset in assets: - value = self.get_single_spot_value( - asset, field, data_frequency) - values.append(value) + values = [] + for asset in assets: + value = self.get_single_spot_value(asset, field, data_frequency) + values.append(value) - return values - else: - return self.get_single_spot_value( - assets, field, data_frequency) + return values def get_single_spot_value(self, asset, field, data_frequency): """ diff --git a/tests/exchange/test_data_portal.py b/tests/exchange/test_data_portal.py index 4e954acf..2b881541 100644 --- a/tests/exchange/test_data_portal.py +++ b/tests/exchange/test_data_portal.py @@ -107,6 +107,6 @@ class ExchangeDataPortalTestCase: date = pd.to_datetime('2017-09-10', utc=True) value = self.data_portal_backtest.get_spot_value( - assets[0], 'close', date, 'minute') + assets, 'close', date, 'minute') pass