Refinements and documentation.

This commit is contained in:
fredfortier
2017-09-23 04:49:13 -04:00
parent bc65c10fc6
commit d6996b1e93
3 changed files with 66 additions and 61 deletions
+59 -49
View File
@@ -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
+6 -11
View File
@@ -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):
"""
+1 -1
View File
@@ -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