# pylint: disable=too-many-locals from typing import Optional import pytest from sqlalchemy import Column, MetaData, Table from sqlalchemy.engine import reflection from dapd_db_schema.schemas.uow_meta import AuditLog, UnitOfWork pytestmark = [pytest.mark.integration] def convert_model_column_type(model_column: Column) -> str: """ Converts model column type (SQLAlchemy) to compare it with reflected one (Postgres). """ type_to_converter_mapping = { 'DATETIME': lambda t: 'TIMESTAMP WITH TIME ZONE' if getattr(t, 'timezone', False) else 'TIMESTAMP', } default_converter = str converter = type_to_converter_mapping.get(str(model_column.type), default_converter) return converter(model_column.type) def get_model_column_server_default(model_column: Column) -> Optional[str]: if model_column.server_default is None: return None return model_column.server_default.arg.text def get_reflected_constrained_columns_length(fk_constraints) -> int: result = 0 for fk_constraint in fk_constraints: result += len(fk_constraint['constrained_columns']) return result @pytest.mark.parametrize( 'model', [ AuditLog, UnitOfWork, ] ) def test_inspect_db(db, model): schema = 'uow_meta' meta = MetaData(schema=schema) meta.reflect(bind=db.uow_meta_engine) model_table: Table = model.metadata.tables[f'{schema}.{model.__tablename__}'] inspector = reflection.Inspector.from_engine(db.uow_meta_engine) # Asserting table name table_names = inspector.get_table_names() assert model_table.fullname in [f'{schema}.{t}' for t in table_names] # Asserting columns and their types for reflected_column in inspector.get_columns(model.__tablename__): column_name = reflected_column['name'] assert hasattr(model_table.columns, column_name), \ f'"{column_name}" was not defined for model' model_column = getattr(model_table.columns, column_name) assert convert_model_column_type(model_column) == str(reflected_column['type']) assert model_column.nullable == reflected_column['nullable'], \ f'"{column_name}" has improper nullable' assert get_model_column_server_default(model_column) == reflected_column['default'], \ f'"{column_name}" has improper default' # Asserting PK constraints pk_constraints = inspector.get_pk_constraint(model.__tablename__) constrained_columns = pk_constraints['constrained_columns'] assert len(model_table.primary_key.columns) == len(constrained_columns), \ f'PK count mismatch for {model_table.name}' for constrained_column in constrained_columns: assert hasattr(model_table.primary_key.columns, constrained_column) # Asserting FK constraints fk_constraints = inspector.get_foreign_keys(model.__tablename__) assert len(model_table.foreign_keys) == get_reflected_constrained_columns_length(fk_constraints) for model_fk in model_table.foreign_keys: for fk_constraint in fk_constraints: if model_fk.constraint.referred_table.fullname == fk_constraint['referred_table']: assert sorted(model_fk.constraint.column_keys) == \ sorted(fk_constraint['constrained_columns']) assert model_fk.column.name in fk_constraint['referred_columns'] assert model_fk.constraint.onupdate == fk_constraint['options']['onupdate'] assert model_fk.constraint.ondelete == fk_constraint['options']['ondelete'] assert model_fk.constraint.deferrable == fk_constraint['options']['deferrable'] assert model_fk.constraint.initially == fk_constraint['options']['initially'] assert model_fk.constraint.match == fk_constraint['options']['match']