diff --git a/tests/test_assets.py b/tests/test_assets.py index e3a42b7d..54ef356d 100644 --- a/tests/test_assets.py +++ b/tests/test_assets.py @@ -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): diff --git a/zipline/assets/assets.py b/zipline/assets/assets.py index 94bcb210..a23c32fd 100644 --- a/zipline/assets/assets.py +++ b/zipline/assets/assets.py @@ -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 diff --git a/zipline/errors.py b/zipline/errors.py index 91e68075..5c7972e2 100644 --- a/zipline/errors.py +++ b/zipline/errors.py @@ -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): diff --git a/zipline/utils/control_flow.py b/zipline/utils/control_flow.py index 891c5ff3..24fe3fcb 100644 --- a/zipline/utils/control_flow.py +++ b/zipline/utils/control_flow.py @@ -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