ENH: Make retrieve specific type functions public.

We rely on these upstream, for better or worse, so add tests and docs.

Also adds distinct `EquitiesNotFound` and `FutureContractsNotFound`
exceptions.
This commit is contained in:
Scott Sanderson
2015-11-13 18:26:54 -05:00
parent 3619a24e4d
commit 657a132f1e
4 changed files with 246 additions and 27 deletions
+159 -3
View File
@@ -28,6 +28,7 @@ import pandas as pd
from pandas.tseries.tools import normalize_date
from pandas.util.testing import assert_frame_equal
from nose_parameterized import parameterized
from numpy import full
import sqlalchemy as sa
@@ -38,6 +39,9 @@ from zipline.assets import (
AssetFinder,
AssetFinderCachedEquities,
)
from six import itervalues
from toolz import valmap
from zipline.assets.futures import (
cme_code_to_month,
FutureChain,
@@ -50,17 +54,23 @@ from zipline.assets.asset_writer import (
_version_table_schema,
)
from zipline.errors import (
SymbolNotFound,
EquitiesNotFound,
FutureContractsNotFound,
MultipleSymbolsFound,
SidAssignmentError,
RootSymbolNotFound,
AssetDBVersionError,
SidAssignmentError,
SidsNotFound,
SymbolNotFound,
)
from zipline.finance.trading import TradingEnvironment, noop_load
from zipline.utils.test_utils import (
all_subindices,
tmp_assets_db,
make_commodity_future_info,
make_rotating_equity_info,
make_simple_equity_info,
tmp_assets_db,
tmp_asset_finder,
)
@@ -838,6 +848,152 @@ class AssetFinderTestCase(TestCase):
self.assertTrue(2 in sids)
self.assertTrue(3 in sids)
def test_group_by_type(self):
equities = make_simple_equity_info(
range(5),
start_date=pd.Timestamp('2014-01-01'),
end_date=pd.Timestamp('2015-01-01'),
)
futures = make_commodity_future_info(
first_sid=6,
root_symbols=['CL'],
years=[2014],
)
# Intersecting sid queries, to exercise loading of partially-cached
# results.
queries = [
([0, 1, 3], [6, 7]),
([0, 2, 3], [7, 10]),
(list(equities.index), list(futures.index)),
]
with tmp_asset_finder(equities=equities, futures=futures) as finder:
for equity_sids, future_sids in queries:
results = finder.group_by_type(equity_sids + future_sids)
self.assertEqual(
results,
{'equity': set(equity_sids), 'future': set(future_sids)},
)
@parameterized.expand([
(Equity, 'retrieve_equities', EquitiesNotFound),
(Future, 'retrieve_futures_contracts', FutureContractsNotFound),
])
def test_retrieve_specific_type(self, type_, lookup_name, failure_type):
equities = make_simple_equity_info(
range(5),
start_date=pd.Timestamp('2014-01-01'),
end_date=pd.Timestamp('2015-01-01'),
)
max_equity = equities.index.max()
futures = make_commodity_future_info(
first_sid=max_equity + 1,
root_symbols=['CL'],
years=[2014],
)
equity_sids = [0, 1]
future_sids = [max_equity + 1, max_equity + 2, max_equity + 3]
if type_ == Equity:
success_sids = equity_sids
fail_sids = future_sids
else:
fail_sids = equity_sids
success_sids = future_sids
with tmp_asset_finder(equities=equities, futures=futures) as finder:
# Run twice to exercise caching.
lookup = getattr(finder, lookup_name)
for _ in range(2):
results = lookup(success_sids)
self.assertIsInstance(results, dict)
self.assertEqual(set(results.keys()), set(success_sids))
self.assertEqual(
valmap(int, results),
dict(zip(success_sids, success_sids)),
)
self.assertEqual(
{type_},
{type(asset) for asset in itervalues(results)},
)
with self.assertRaises(failure_type):
lookup(fail_sids)
with self.assertRaises(failure_type):
# Should fail if **any** of the assets are bad.
lookup([success_sids[0], fail_sids[0]])
def test_retrieve_all(self):
equities = make_simple_equity_info(
range(5),
start_date=pd.Timestamp('2014-01-01'),
end_date=pd.Timestamp('2015-01-01'),
)
max_equity = equities.index.max()
futures = make_commodity_future_info(
first_sid=max_equity + 1,
root_symbols=['CL'],
years=[2014],
)
with tmp_asset_finder(equities=equities, futures=futures) as finder:
all_sids = finder.sids
self.assertEqual(len(all_sids), len(equities) + len(futures))
queries = [
# Empty Query.
(),
# Only Equities.
tuple(equities.index[:2]),
# Only Futures.
tuple(futures.index[:3]),
# Mixed, all cache misses.
tuple(equities.index[2:]) + tuple(futures.index[3:]),
# Mixed, all cache hits.
tuple(equities.index[2:]) + tuple(futures.index[3:]),
# Everything.
all_sids,
all_sids,
]
for sids in queries:
equity_sids = [i for i in sids if i <= max_equity]
future_sids = [i for i in sids if i > max_equity]
results = finder.retrieve_all(sids)
self.assertEqual(sids, tuple(map(int, results)))
self.assertEqual(
[Equity for _ in equity_sids] +
[Future for _ in future_sids],
list(map(type, results)),
)
self.assertEqual(
(
list(equities.symbol.loc[equity_sids]) +
list(futures.symbol.loc[future_sids])
),
list(asset.symbol for asset in results),
)
@parameterized.expand([
(EquitiesNotFound, 'equity', 'equities'),
(FutureContractsNotFound, 'future contract', 'future contracts'),
(SidsNotFound, 'asset', 'assets'),
])
def test_error_message_plurality(self,
error_type,
singular,
plural):
try:
raise error_type(sids=[1])
except error_type as e:
self.assertEqual(
str(e),
"No {singular} found for sid: 1.".format(singular=singular)
)
try:
raise error_type(sids=[1, 2])
except error_type as e:
self.assertEqual(
str(e),
"No {plural} found for sids: [1, 2].".format(plural=plural)
)
class AssetFinderCachedEquitiesTestCase(AssetFinderTestCase):
+50 -16
View File
@@ -26,11 +26,13 @@ from six.moves import map as imap
import sqlalchemy as sa
from zipline.errors import (
EquitiesNotFound,
FutureContractsNotFound,
MapAssetIdentifierIndexError,
MultipleSymbolsFound,
RootSymbolNotFound,
SidsNotFound,
SymbolNotFound,
MapAssetIdentifierIndexError,
)
from zipline.assets import (
Asset, Equity, Future,
@@ -227,10 +229,10 @@ class AssetFinder(object):
raise SidsNotFound(sids=list(failures))
# We don't update the asset cache here because it should already be
# updated by `self._retrieve_equities`.
update_hits(self._retrieve_equities(type_to_assets.pop('equity', ())))
# updated by `self.retrieve_equities`.
update_hits(self.retrieve_equities(type_to_assets.pop('equity', ())))
update_hits(
self._retrieve_futures_contracts(type_to_assets.pop('future', ()))
self.retrieve_futures_contracts(type_to_assets.pop('future', ()))
)
# We shouldn't know about any other asset types.
@@ -241,18 +243,52 @@ class AssetFinder(object):
return [hits[sid] for sid in sids]
def _retrieve_equities(self, sids):
def retrieve_equities(self, sids):
"""
Retrieve the Equity object of a given sid.
Retrieve Equity objects for a list of sids.
Users generally shouldn't need to this method (instead, they should
prefer the more general/friendly `retrieve_assets`), but it has a
documented interface and tests because it's used upstream.
Parameters
----------
sids : iterable[int]
Returns
-------
equities : dict[int -> Equity]
Raises
------
EquitiesNotFound
When any requested asset isn't found.
"""
return self._retrieve_assets(sids, self.equities, Equity)
def _retrieve_equity(self, sid):
return self._retrieve_equities((sid,))[sid]
return self.retrieve_equities((sid,))[sid]
def _retrieve_futures_contracts(self, sids):
def retrieve_futures_contracts(self, sids):
"""
Retrieve the Future object of a given sid.
Retrieve Future objects for an iterable of sids.
Users generally shouldn't need to this method (instead, they should
prefer the more general/friendly `retrieve_assets`), but it has a
documented interface and tests because it's used upstream.
Parameters
----------
sids : iterable[int]
Returns
-------
equities : dict[int -> Equity]
Raises
------
EquitiesNotFound
When any requested asset isn't found.
"""
return self._retrieve_assets(sids, self.futures_contracts, Future)
@@ -307,12 +343,10 @@ class AssetFinder(object):
# an error in our code, not a user-input error.
misses = tuple(set(sids) - viewkeys(hits))
if misses:
raise AssertionError(
"Couldn't resolve sids {sids} as instances of {type}.".format(
sids=misses,
type=asset_type,
)
)
if asset_type == Equity:
raise EquitiesNotFound(sids=misses)
else:
raise FutureContractsNotFound(sids=misses)
return hits
def _get_fuzzy_candidates(self, fuzzy_symbol):
@@ -589,7 +623,7 @@ class AssetFinder(object):
if count == 0:
raise RootSymbolNotFound(root_symbol=root_symbol)
contracts = self._retrieve_futures_contracts(sids)
contracts = self.retrieve_futures_contracts(sids)
return [contracts[sid] for sid in sids]
@property
+33 -4
View File
@@ -237,12 +237,41 @@ class SidsNotFound(ZiplineError):
Raised when a retrieve_asset() or retrieve_all() call contains a
non-existent sid.
"""
@lazyval
def plural(self):
return len(self.sids) > 1
@lazyval
def sids(self):
return self.kwargs['sids']
@lazyval
def msg(self):
sids = self.kwargs['sids']
if len(sids) == 1:
return "No asset found for sid: {sids[0]}."
return "No assets found for sids: {sids}."
if self.plural:
return "No assets found for sids: {sids}."
return "No asset found for sid: {sids[0]}."
class EquitiesNotFound(SidsNotFound):
"""
Raised when a call to `retrieve_equities` fails to find an asset.
"""
@lazyval
def msg(self):
if self.plural:
return "No equities found for sids: {sids}."
return "No equity found for sid: {sids[0]}."
class FutureContractsNotFound(SidsNotFound):
"""
Raised when a call to `retrieve_futures_contracts` fails to find an asset.
"""
@lazyval
def msg(self):
if self.plural:
return "No future contracts found for sids: {sids}."
return "No future contract found for sid: {sids[0]}."
class ConsumeAssetMetaDataError(ZiplineError):
+4 -4
View File
@@ -59,15 +59,15 @@ def ignore_nanwarnings():
def invert(d):
"""
Invert a dictionary into a dictionary of lists.
Invert a dictionary into a dictionary of sets.
>>> invert({'a': 1, 'b': 2, 'c': 1})
{1: ['a', 'c'], 2: ['b']}
{1: {'a', 'c'}, 2: {'b'}}
"""
out = {}
for k, v in iteritems(d):
try:
out[v].append(k)
out[v].add(k)
except KeyError:
out[v] = [k]
out[v] = {k}
return out