Merge pull request #1387 from quantopian/sqlite-needs-distinct-on-because-this-is-ugly

v5 downgrade
This commit is contained in:
Joe Jevnik
2016-08-10 16:10:04 -04:00
committed by GitHub
3 changed files with 94 additions and 27 deletions
+33
View File
@@ -1470,3 +1470,36 @@ class TestAssetDBVersioning(ZiplineTestCase):
# higher-than-current version
with self.assertRaises(AssetDBImpossibleDowngrade):
downgrade(self.engine, ASSET_DB_VERSION + 5)
def test_v5_to_v4_selects_most_recent_ticker(self):
T = pd.Timestamp
AssetDBWriter(self.engine).write(
equities=pd.DataFrame(
[['A', 'A', T('2014-01-01'), T('2014-01-02')],
['B', 'B', T('2014-01-01'), T('2014-01-02')],
# these two are both ticker sid 2
['B', 'C', T('2014-01-03'), T('2014-01-04')],
['C', 'C', T('2014-01-01'), T('2014-01-02')]],
index=[0, 1, 2, 2],
columns=['symbol', 'asset_name', 'start_date', 'end_date'],
),
)
downgrade(self.engine, 4)
metadata = sa.MetaData(self.engine)
metadata.reflect()
def select_fields(r):
return r.sid, r.symbol, r.asset_name, r.start_date, r.end_date
expected_data = {
(0, 'A', 'A', T('2014-01-01').value, T('2014-01-02').value),
(1, 'B', 'B', T('2014-01-01').value, T('2014-01-02').value),
(2, 'B', 'C', T('2014-01-01').value, T('2014-01-04').value),
}
actual_data = set(map(
select_fields,
sa.select(metadata.tables['equities'].c).execute(),
))
assert_equal(expected_data, actual_data)
+12 -4
View File
@@ -272,14 +272,22 @@ def _downgrade_v5(op):
from
equities
inner join
-- Nested select here to take the most recently held ticker
-- for each sid. The group by with no aggregation function will
-- take the last element in the group, so we first order by
-- the end date ascending to ensure that the groupby takes
-- the last ticker.
(select
*
from
equity_symbol_mappings
(select
*
from
equity_symbol_mappings
order by
equity_symbol_mappings.end_date asc)
group by
equity_symbol_mappings.sid
order by
equity_symbol_mappings.end_date desc) sym
sid) sym
on
equities.sid == sym.sid
""",
+49 -23
View File
@@ -30,7 +30,6 @@ from nose.tools import ( # noqa
assert_raises_regexp,
assert_regexp_matches,
assert_sequence_equal,
assert_set_equal,
assert_true,
assert_tuple_equal,
)
@@ -302,37 +301,54 @@ def assert_float_equal(result,
)
@assert_equal.register(dict, dict)
def assert_dict_equal(result, expected, path=(), msg='', **kwargs):
if path is None:
path = ()
def _check_sets(result, expected, msg, path, type_):
"""Compare two sets. This is used to check dictionary keys and sets.
result_keys = viewkeys(result)
expected_keys = viewkeys(expected)
if result_keys != expected_keys:
if result_keys > expected_keys:
diff = result_keys - expected_keys
msg = 'extra %s in result: %r' % (_s('key', diff), diff)
elif result_keys < expected_keys:
diff = expected_keys - result_keys
msg = 'result is missing %s: %r' % (_s('key', diff), diff)
Parameters
----------
result : set
expected : set
msg : str
path : tuple
type : str
The type of an element. For dict we use ``'key'`` and for set we use
``'element'``.
"""
if result != expected:
if result > expected:
diff = result - expected
msg = 'extra %s in result: %r' % (_s(type_, diff), diff)
elif result < expected:
diff = expected - result
msg = 'result is missing %s: %r' % (_s(type_, diff), diff)
else:
sym = result_keys ^ expected_keys
in_result = sym - expected_keys
in_expected = sym - result_keys
in_result = result - expected
in_expected = expected - result
msg = '%s only in result: %s\n%s only in expected: %s' % (
_s('key', in_result),
_s(type_, in_result),
in_result,
_s('key', in_expected),
_s(type_, in_expected),
in_expected,
)
raise AssertionError(
'%sdict keys do not match\n%s' % (
'%s%ss do not match\n%s' % (
_fmt_msg(msg),
_fmt_path(path + ('.%s()' % ('viewkeys' if PY2 else 'keys'),)),
type_,
_fmt_path(path),
),
)
@assert_equal.register(dict, dict)
def assert_dict_equal(result, expected, path=(), msg='', **kwargs):
_check_sets(
viewkeys(result),
viewkeys(expected),
msg,
path + ('.%s()' % ('viewkeys' if PY2 else 'keys'),),
'key',
)
failures = []
for k, (resultv, expectedv) in iteritems(dzip_exact(result, expected)):
try:
@@ -350,7 +366,7 @@ def assert_dict_equal(result, expected, path=(), msg='', **kwargs):
raise AssertionError('\n'.join(failures))
@assert_equal.register(list, list) # noqa
@assert_equal.register(list, list)
def assert_list_equal(result, expected, path=(), msg='', **kwargs):
result_len = len(result)
expected_len = len(expected)
@@ -362,7 +378,6 @@ def assert_list_equal(result, expected, path=(), msg='', **kwargs):
_fmt_path(path),
)
)
for n, (resultv, expectedv) in enumerate(zip(result, expected)):
assert_equal(
resultv,
@@ -373,6 +388,17 @@ def assert_list_equal(result, expected, path=(), msg='', **kwargs):
)
@assert_equal.register(set, set)
def assert_set_equal(result, expected, path=(), msg='', **kwargs):
_check_sets(
result,
expected,
msg,
path,
'element',
)
@assert_equal.register(np.ndarray, np.ndarray)
def assert_array_equal(result,
expected,