diff --git a/tests/test_assets.py b/tests/test_assets.py index 85a170c6..11b92d31 100644 --- a/tests/test_assets.py +++ b/tests/test_assets.py @@ -50,6 +50,8 @@ from zipline.assets.futures import ( from zipline.assets.asset_writer import ( check_version_info, write_version_info, +) +from zipline.assets.asset_db_schema import ( ASSET_DB_VERSION, _version_table_schema, ) @@ -1331,3 +1333,24 @@ class TestAssetDBVersioning(TestCase): # Assert that trying to overwrite the version fails with self.assertRaises(sa.exc.IntegrityError): write_version_info(version_table, -3) + + def test_finder_checks_version(self): + # Create an env and give it a bogus version number + env = TradingEnvironment(load=noop_load) + metadata = sa.MetaData(bind=env.engine) + version_table = _version_table_schema(metadata) + version_table.delete().execute() + write_version_info(version_table, -2) + check_version_info(version_table, -2) + + # Assert that trying to build a finder with a bad db raises an error + with self.assertRaises(AssetDBVersionError): + AssetFinder(engine=env.engine) + + # Change the version number of the db to the correct version + version_table.delete().execute() + write_version_info(version_table, ASSET_DB_VERSION) + check_version_info(version_table, ASSET_DB_VERSION) + + # Now that the versions match, this Finder should succeed + AssetFinder(engine=env.engine) diff --git a/zipline/assets/asset_db_schema.py b/zipline/assets/asset_db_schema.py new file mode 100644 index 00000000..29f6b568 --- /dev/null +++ b/zipline/assets/asset_db_schema.py @@ -0,0 +1,162 @@ +import sqlalchemy as sa + + +# Define a version number for the database generated by these writers +# Increment this version number any time a change is made to the schema of the +# assets database +ASSET_DB_VERSION = 0 + + +def generate_asset_db_metadata(bind=None): + # NOTE: When modifying this schema, update the ASSET_DB_VERSION value + metadata = sa.MetaData(bind=bind) + _version_table_schema(metadata) + _equities_table_schema(metadata) + _futures_exchanges_schema(metadata) + _futures_root_symbols_schema(metadata) + _futures_contracts_schema(metadata) + _asset_router_schema(metadata) + return metadata + + +# A list of the names of all tables in the assets db +# NOTE: When modifying this schema, update the ASSET_DB_VERSION value +asset_db_table_names = ['version_info', 'equities', 'futures_exchanges', + 'futures_root_symbols', 'futures_contracts', + 'asset_router'] + + +def _equities_table_schema(metadata): + # NOTE: When modifying this schema, update the ASSET_DB_VERSION value + return sa.Table( + 'equities', + metadata, + sa.Column( + 'sid', + sa.Integer, + unique=True, + nullable=False, + primary_key=True, + ), + sa.Column('symbol', sa.Text), + sa.Column('company_symbol', sa.Text, index=True), + sa.Column('share_class_symbol', sa.Text), + sa.Column('fuzzy_symbol', sa.Text, index=True), + sa.Column('asset_name', sa.Text), + sa.Column('start_date', sa.Integer, default=0, nullable=False), + sa.Column('end_date', sa.Integer, nullable=False), + sa.Column('first_traded', sa.Integer, nullable=False), + sa.Column('exchange', sa.Text), + ) + + +def _futures_exchanges_schema(metadata): + # NOTE: When modifying this schema, update the ASSET_DB_VERSION value + return sa.Table( + 'futures_exchanges', + metadata, + sa.Column( + 'exchange', + sa.Text, + unique=True, + nullable=False, + primary_key=True, + ), + sa.Column('timezone', sa.Text), + ) + + +def _futures_root_symbols_schema(metadata): + # NOTE: When modifying this schema, update the ASSET_DB_VERSION value + return sa.Table( + 'futures_root_symbols', + metadata, + sa.Column( + 'root_symbol', + sa.Text, + unique=True, + nullable=False, + primary_key=True, + ), + sa.Column('root_symbol_id', sa.Integer), + sa.Column('sector', sa.Text), + sa.Column('description', sa.Text), + sa.Column( + 'exchange', + sa.Text, + sa.ForeignKey('futures_exchanges.exchange'), + ), + ) + + +def _futures_contracts_schema(metadata): + # NOTE: When modifying this schema, update the ASSET_DB_VERSION value + return sa.Table( + 'futures_contracts', + metadata, + sa.Column( + 'sid', + sa.Integer, + unique=True, + nullable=False, + primary_key=True, + ), + sa.Column('symbol', sa.Text, unique=True, index=True), + sa.Column( + 'root_symbol', + sa.Text, + sa.ForeignKey('futures_root_symbols.root_symbol'), + index=True + ), + sa.Column('asset_name', sa.Text), + sa.Column('start_date', sa.Integer, default=0, nullable=False), + sa.Column('end_date', sa.Integer, nullable=False), + sa.Column('first_traded', sa.Integer, nullable=False), + sa.Column( + 'exchange', + sa.Text, + sa.ForeignKey('futures_exchanges.exchange'), + ), + sa.Column('notice_date', sa.Integer, nullable=False), + sa.Column('expiration_date', sa.Integer, nullable=False), + sa.Column('auto_close_date', sa.Integer, nullable=False), + sa.Column('contract_multiplier', sa.Float), + ) + + +def _asset_router_schema(metadata): + # NOTE: When modifying this schema, update the ASSET_DB_VERSION value + return sa.Table( + 'asset_router', + metadata, + sa.Column( + 'sid', + sa.Integer, + unique=True, + nullable=False, + primary_key=True), + sa.Column('asset_type', sa.Text), + ) + + +def _version_table_schema(metadata): + # NOTE: When modifying this schema, update the ASSET_DB_VERSION value + return sa.Table( + 'version_info', + metadata, + sa.Column( + 'id', + sa.Integer, + unique=True, + nullable=False, + primary_key=True, + ), + sa.Column( + 'version', + sa.Integer, + unique=True, + nullable=False, + ), + # This constraint ensures a single entry in this table + sa.CheckConstraint('id <= 1'), + ) diff --git a/zipline/assets/asset_writer.py b/zipline/assets/asset_writer.py index 328b02e1..593a0ea0 100644 --- a/zipline/assets/asset_writer.py +++ b/zipline/assets/asset_writer.py @@ -27,6 +27,11 @@ import sqlalchemy as sa from zipline.errors import SidAssignmentError, AssetDBVersionError from zipline.assets._assets import Asset +from zipline.assets.asset_db_schema import ( + generate_asset_db_metadata, + asset_db_table_names, + ASSET_DB_VERSION, +) SQLITE_MAX_VARIABLE_NUMBER = 999 @@ -164,12 +169,6 @@ def _generate_output_dataframe(data_subset, defaults): return output -# Define a version number for the database generated by these writers -# Increment this version number any time a breaking change is made to the -# schema and readers of the database -ASSET_DB_VERSION = 0 - - def check_version_info(version_table, expected_version): """ Checks for a version value in the version table. @@ -215,143 +214,6 @@ def write_version_info(version_table, version_value): sa.insert(version_table, values={'version': version_value}).execute() -# A list of the names of all tables in the assets db -asset_db_table_names = ['version_info', 'equities', 'futures_exchanges', - 'futures_root_symbols', 'futures_contracts', - 'asset_router'] - - -def _equities_table_schema(metadata): - return sa.Table( - 'equities', - metadata, - sa.Column( - 'sid', - sa.Integer, - unique=True, - nullable=False, - primary_key=True, - ), - sa.Column('symbol', sa.Text), - sa.Column('company_symbol', sa.Text, index=True), - sa.Column('share_class_symbol', sa.Text), - sa.Column('fuzzy_symbol', sa.Text, index=True), - sa.Column('asset_name', sa.Text), - sa.Column('start_date', sa.Integer, default=0, nullable=False), - sa.Column('end_date', sa.Integer, nullable=False), - sa.Column('first_traded', sa.Integer, nullable=False), - sa.Column('exchange', sa.Text), - ) - - -def _futures_exchanges_schema(metadata): - return sa.Table( - 'futures_exchanges', - metadata, - sa.Column( - 'exchange', - sa.Text, - unique=True, - nullable=False, - primary_key=True, - ), - sa.Column('timezone', sa.Text), - ) - - -def _futures_root_symbols_schema(metadata, futures_exchanges): - return sa.Table( - 'futures_root_symbols', - metadata, - sa.Column( - 'root_symbol', - sa.Text, - unique=True, - nullable=False, - primary_key=True, - ), - sa.Column('root_symbol_id', sa.Integer), - sa.Column('sector', sa.Text), - sa.Column('description', sa.Text), - sa.Column( - 'exchange', - sa.Text, - sa.ForeignKey(futures_exchanges.c.exchange), - ), - ) - - -def _futures_contracts_schema(metadata, futures_root_symbols, - futures_exchanges): - return sa.Table( - 'futures_contracts', - metadata, - sa.Column( - 'sid', - sa.Integer, - unique=True, - nullable=False, - primary_key=True, - ), - sa.Column('symbol', sa.Text, unique=True, index=True), - sa.Column( - 'root_symbol', - sa.Text, - sa.ForeignKey(futures_root_symbols.c.root_symbol), - index=True - ), - sa.Column('asset_name', sa.Text), - sa.Column('start_date', sa.Integer, default=0, nullable=False), - sa.Column('end_date', sa.Integer, nullable=False), - sa.Column('first_traded', sa.Integer, nullable=False), - sa.Column( - 'exchange', - sa.Text, - sa.ForeignKey(futures_exchanges.c.exchange), - ), - sa.Column('notice_date', sa.Integer, nullable=False), - sa.Column('expiration_date', sa.Integer, nullable=False), - sa.Column('auto_close_date', sa.Integer, nullable=False), - sa.Column('contract_multiplier', sa.Float), - ) - - -def _asset_router_schema(metadata): - return sa.Table( - 'asset_router', - metadata, - sa.Column( - 'sid', - sa.Integer, - unique=True, - nullable=False, - primary_key=True), - sa.Column('asset_type', sa.Text), - ) - - -def _version_table_schema(metadata): - return sa.Table( - 'version_info', - metadata, - sa.Column( - 'id', - sa.Integer, - unique=True, - nullable=False, - primary_key=True, - ), - sa.Column( - 'version', - sa.Integer, - unique=True, - nullable=False, - ), - # This constraint ensures a single entry in this table - sa.CheckConstraint('id <= 1'), - ) - - class AssetDBWriter(with_metaclass(ABCMeta)): """ Class used to write arbitrary data to SQLite database. @@ -480,23 +342,12 @@ class AssetDBWriter(with_metaclass(ABCMeta)): constraints : bool, optional If True, create SQL ForeignKey and PrimaryKey constraints. """ - metadata = sa.MetaData(bind=engine) + tables_already_exist = self.check_for_tables(engine) + metadata = generate_asset_db_metadata(bind=engine) - tables_already_exist = self.check_for_tables(engine=engine) + for table_name in asset_db_table_names: + setattr(self, table_name, metadata.tables[table_name]) - self.equities = _equities_table_schema(metadata) - self.futures_exchanges = _futures_exchanges_schema(metadata) - self.futures_root_symbols = _futures_root_symbols_schema( - metadata=metadata, - futures_exchanges=self.futures_exchanges, - ) - self.futures_contracts = _futures_contracts_schema( - metadata=metadata, - futures_root_symbols=self.futures_root_symbols, - futures_exchanges=self.futures_exchanges, - ) - self.asset_router = _asset_router_schema(metadata) - self.version_info = _version_table_schema(metadata) # Create the SQL tables if they do not already exist. metadata.create_all(checkfirst=True) diff --git a/zipline/assets/assets.py b/zipline/assets/assets.py index 712c24de..13b457b4 100644 --- a/zipline/assets/assets.py +++ b/zipline/assets/assets.py @@ -37,11 +37,13 @@ from zipline.assets import ( Asset, Equity, Future, ) from zipline.assets.asset_writer import ( - split_delimited_symbol, check_version_info, - ASSET_DB_VERSION, + split_delimited_symbol, asset_db_table_names, - SQLITE_MAX_VARIABLE_NUMBER + SQLITE_MAX_VARIABLE_NUMBER, +) +from zipline.assets.asset_db_schema import ( + ASSET_DB_VERSION ) from zipline.utils.control_flow import invert