From 1f10fff1c4a2d94863c8e9ffac54f5001785c9ad Mon Sep 17 00:00:00 2001 From: Joe Jevnik Date: Tue, 2 Aug 2016 18:34:00 -0400 Subject: [PATCH] BUG: support querying more than 999 assets at a time --- tests/test_assets.py | 26 +++++++++++++++++++++++- zipline/assets/assets.py | 43 +++++++++++++++++++++++++++------------- 2 files changed, 54 insertions(+), 15 deletions(-) diff --git a/tests/test_assets.py b/tests/test_assets.py index b81fac57..13f19e75 100644 --- a/tests/test_assets.py +++ b/tests/test_assets.py @@ -30,7 +30,7 @@ from nose_parameterized import parameterized from numpy import full, int32, int64 import pandas as pd from pandas.util.testing import assert_frame_equal -from six import PY2 +from six import PY2, viewkeys import sqlalchemy as sa from zipline.assets import ( @@ -57,6 +57,7 @@ from zipline.assets.asset_writer import ( check_version_info, write_version_info, _futures_defaults, + SQLITE_MAX_VARIABLE_NUMBER, ) from zipline.assets.asset_db_schema import ASSET_DB_VERSION from zipline.assets.asset_db_migrations import ( @@ -83,6 +84,7 @@ from zipline.testing.fixtures import ( ZiplineTestCase, WithTradingCalendar, ) +from zipline.utils.range import range @contextmanager @@ -407,6 +409,28 @@ class AssetFinderTestCase(WithTradingCalendar, ZiplineTestCase): self._asset_writer = AssetDBWriter(conn) self.asset_finder = self.asset_finder_type(conn) + def test_blocked_lookup_symbol_query(self): + # we will try to query for more variables than sqlite supports + # to make sure we are properly chunking on the client side + as_of = pd.Timestamp('2013-01-01', tz='UTC') + # we need more sids than we can query from sqlite + nsids = SQLITE_MAX_VARIABLE_NUMBER + 10 + sids = range(nsids) + frame = pd.DataFrame.from_records( + [ + { + 'sid': sid, + 'symbol': 'TEST.%d' % sid, + 'start_date': as_of.value, + 'end_date': as_of.value, + } + for sid in sids + ] + ) + self.write_assets(equities=frame) + assets = self.asset_finder.retrieve_equities(sids) + assert_equal(viewkeys(assets), set(sids)) + def test_lookup_symbol_delimited(self): as_of = pd.Timestamp('2013-01-01', tz='UTC') frame = pd.DataFrame.from_records( diff --git a/zipline/assets/assets.py b/zipline/assets/assets.py index 4796ac8d..b770014a 100644 --- a/zipline/assets/assets.py +++ b/zipline/assets/assets.py @@ -23,7 +23,16 @@ import pandas as pd from pandas import isnull from six import with_metaclass, string_types, viewkeys, iteritems import sqlalchemy as sa -from toolz import merge, compose, valmap, sliding_window, concatv, curry +from toolz import ( + compose, + concat, + concatv, + curry, + merge, + partition_all, + sliding_window, + valmap, +) from toolz.curried import operator as op from zipline.errors import ( @@ -43,6 +52,7 @@ from .asset_writer import ( split_delimited_symbol, asset_db_table_names, symbol_columns, + SQLITE_MAX_VARIABLE_NUMBER, ) from .asset_db_schema import ( ASSET_DB_VERSION @@ -432,21 +442,26 @@ class AssetFinder(object): def _lookup_most_recent_symbols(self, sids): symbol_cols = self.equity_symbol_mappings.c - symbols = { row.sid: {c: row[c] for c in symbol_columns} - for row in self.engine.execute( - sa.select( - (symbol_cols.sid,) + - tuple(map(op.getitem(symbol_cols), symbol_columns)), - ).where( - symbol_cols.sid.in_(map(int, sids)), - ).order_by( - symbol_cols.end_date.desc(), - ).group_by( - symbol_cols.sid, - ) - ).fetchall() + for row in concat( + self.engine.execute( + sa.select( + (symbol_cols.sid,) + + tuple(map(op.getitem(symbol_cols), symbol_columns)), + ).where( + symbol_cols.sid.in_(map(int, sid_group)), + ).order_by( + symbol_cols.end_date.desc(), + ).group_by( + symbol_cols.sid, + ) + ).fetchall() + for sid_group in partition_all( + SQLITE_MAX_VARIABLE_NUMBER, + sids + ), + ) } if len(symbols) != len(sids):