BUG: support querying more than 999 assets at a time

This commit is contained in:
Joe Jevnik
2016-08-02 18:53:57 -04:00
parent 7e99094cb1
commit 1f10fff1c4
2 changed files with 54 additions and 15 deletions
+25 -1
View File
@@ -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(
+29 -14
View File
@@ -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):