diff --git a/ibm_db_sa/base.py b/ibm_db_sa/base.py index b0169ce..0149445 100644 --- a/ibm_db_sa/base.py +++ b/ibm_db_sa/base.py @@ -21,7 +21,7 @@ """ import sys import sqlalchemy -import datetime, re +import datetime, decimal, re from sqlalchemy import types as sa_types from sqlalchemy import schema as sa_schema from sqlalchemy import util @@ -228,6 +228,30 @@ class DOUBLE(sa_types.Numeric): __visit_name__ = 'DOUBLE' +class DECFLOAT(sa_types.Numeric): + """DB2 DECFLOAT(16) or DECFLOAT(34). + + The ibm_db DBAPI returns DECFLOAT values as str; convert them to Decimal + (or float with asdecimal=False) and bind Decimal values as exact text. + """ + __visit_name__ = 'DECFLOAT' + + def __init__(self, precision=34, asdecimal=True): + super().__init__(precision=precision, asdecimal=asdecimal) + + def bind_processor(self, dialect): + def process(value): + return str(value) if isinstance(value, decimal.Decimal) else value + return process + + def result_processor(self, dialect, coltype): + convert = decimal.Decimal if self.asdecimal else float + + def process(value): + return None if value is None else convert(str(value)) + return process + + class LONGVARCHAR(sa_types.VARCHAR): __visit_name_ = 'LONGVARCHAR' @@ -283,6 +307,9 @@ class XML(sa_types.Text): 'XML': XML, 'GRAPHIC': GRAPHIC, 'VARGRAPHIC': VARGRAPHIC, + 'DECFLOAT': DECFLOAT, + 'BINARY': sa_types.BINARY, + 'VARBINARY': sa_types.VARBINARY, 'LONGVARGRAPHIC': LONGVARGRAPHIC, 'DBCLOB': DBCLOB } @@ -301,6 +328,10 @@ def visit_DATE(self, type_, **kw): logger.debug(f"Type rendering -> DATE -> {sql}") return sql + @log_entry_exit + def visit_DECFLOAT(self, type_, **kw): + return "DECFLOAT(%d)" % (type_.precision or 34) + @log_entry_exit def visit_TIME(self, type_, **kw): sql = "TIME" @@ -1481,6 +1512,13 @@ def get_indexes(self, connection, table_name, schema=None, **kw): logger.debug(f"Indexes fetched -> count={len(indexes)}") return indexes + @log_entry_exit + def get_check_constraints(self, connection, table_name, schema=None, **kw): + reflect = getattr(self._reflector, "get_check_constraints", None) + if reflect is None: + raise NotImplementedError() + return reflect(connection, table_name, schema=schema, **kw) + @log_entry_exit def get_unique_constraints(self, connection, table_name, schema=None, **kw): logger.debug(f"Fetching unique constraints -> table={table_name}, schema={schema}") diff --git a/ibm_db_sa/ibm_db.py b/ibm_db_sa/ibm_db.py index 9ddb9ae..0c1cae4 100644 --- a/ibm_db_sa/ibm_db.py +++ b/ibm_db_sa/ibm_db.py @@ -29,7 +29,7 @@ from sqlalchemy.engine.url import URL from sqlalchemy.exc import ArgumentError -from .base import DB2Dialect, DB2ExecutionContext +from .base import DB2Dialect, DB2ExecutionContext, DECFLOAT, XML from .logger import init_ibmdbsa_logging, log_entry_exit, logger m = re.match(r"^\s*(\d+)\.(\d+)", SA_VERSION_STR) @@ -64,6 +64,47 @@ def to_float(value): return to_float +class _IBM_Binary_ibm_db(sa_types._Binary): + """Bind binary values as bytes. + + SQLAlchemy binds binary values through dbapi.Binary, which ibm_db_dbi + implements as memoryview. With executemany, ibm_db rejects a memoryview for + BINARY, VARBINARY and FOR BIT DATA columns (SQL0302N) and stores its repr + text ("") in BLOB columns. bytes work in both paths. + """ + + def bind_processor(self, dialect): + def process(value): + return None if value is None else bytes(value) + return process + + +# DB2 stores XML parsed, without a declaration. On fetch the CLI serializes it +# with a byte order mark and a UTF-16 declaration, even for a document stored +# with a UTF-8 declaration, which is wrong for a Python str. +_CLI_XML_PREFIX = re.compile(r'\A\ufeff?(?:<\?xml version="1\.0" encoding="UTF-16" \?>)?') + + +class _IBM_XML_ibm_db(XML): + def result_processor(self, dialect, coltype): + def process(value): + if isinstance(value, str): + return _CLI_XML_PREFIX.sub("", value, count=1) + return value + return process + + +_DISCONNECT_MESSAGES = ( + 'Connection is not active', + 'connection is no longer active', + 'Connection Resource cannot be found', + 'SQL30081N', + 'CLI0108E', + 'CLI0106E', + 'SQL1224N', +) + + class DB2ExecutionContext_ibm_db(DB2ExecutionContext): _callproc_result = None _out_parameters = None @@ -128,7 +169,11 @@ class DB2Dialect_ibm_db(DB2Dialect): colspecs = util.update_copy( DB2Dialect.colspecs, { - sa_types.Numeric: _IBM_Numeric_ibm_db + sa_types.Numeric: _IBM_Numeric_ibm_db, + sa_types._Binary: _IBM_Binary_ibm_db, + # DECFLOAT is a Numeric but keeps its own processors. + DECFLOAT: DECFLOAT, + XML: _IBM_XML_ibm_db, } ) @@ -303,27 +348,25 @@ def _get_default_schema_name(self, connection): logger.debug("Normalized schema: %s", normalized_schema_name) return normalized_schema_name - # Checks if the DB_API driver error indicates an invalid connection + # Checks if the DB_API driver error indicates an invalid connection. A + # connection lost while fetching surfaces as the base ibm_db_dbi.Error. @log_entry_exit def is_disconnect(self, ex, connection, cursor): - logger.debug("Checking if exception indicates disconnect") - logger.debug("Exception received: %s", ex) - if isinstance(ex, (self.dbapi.ProgrammingError, - self.dbapi.OperationalError)): - connection_errors = ('Connection is not active', - 'connection is no longer active', - 'Connection Resource cannot be found', - 'SQL30081N', - 'CLI0108E', - 'CLI0106E', - 'SQL1224N') - for err_msg in connection_errors: + if isinstance(ex, self.dbapi.Error): + for err_msg in _DISCONNECT_MESSAGES: if err_msg in str(ex): logger.debug("Disconnect detected due to error: %s", err_msg) return True - else: - logger.debug("Exception type does not indicate disconnect") return False + # After a server restart or network failure ibm_db_dbi raises CLI0106E + # ("Connection is closed") from close(); the pool then logged an error for + # every connection it discarded. + def do_close(self, dbapi_connection): + try: + dbapi_connection.close() + except self.dbapi.Error as err: + if "CLI0106E" not in str(err): + raise dialect = DB2Dialect_ibm_db diff --git a/ibm_db_sa/reflection.py b/ibm_db_sa/reflection.py index 1824a58..e3fca87 100644 --- a/ibm_db_sa/reflection.py +++ b/ibm_db_sa/reflection.py @@ -184,6 +184,35 @@ class DB2Reflector(BaseReflector): Column("COLNAMES", CoerceUnicode, key="colnames"), Column("UNIQUERULE", CoerceUnicode, key="uniquerule"), Column("SYSTEM_REQUIRED", sa_types.SMALLINT, key="system_required"), + Column("INDSCHEMA", CoerceUnicode, key="indschema"), + Column("INDEXTYPE", CoerceUnicode, key="indextype"), + schema="SYSCAT") + + sys_indexcoluse = Table("INDEXCOLUSE", ischema, + Column("INDSCHEMA", CoerceUnicode, key="indschema"), + Column("INDNAME", CoerceUnicode, key="indname"), + Column("COLNAME", CoerceUnicode, key="colname"), + Column("COLSEQ", sa_types.SMALLINT, key="colseq"), + Column("COLORDER", CoerceUnicode, key="colorder"), + schema="SYSCAT") + + sys_references = Table("REFERENCES", ischema, + Column("CONSTNAME", CoerceUnicode, key="constname"), + Column("TABSCHEMA", CoerceUnicode, key="tabschema"), + Column("TABNAME", CoerceUnicode, key="tabname"), + Column("REFKEYNAME", CoerceUnicode, key="refkeyname"), + Column("REFTABSCHEMA", CoerceUnicode, key="reftabschema"), + Column("REFTABNAME", CoerceUnicode, key="reftabname"), + Column("DELETERULE", CoerceUnicode, key="deleterule"), + Column("UPDATERULE", CoerceUnicode, key="updaterule"), + schema="SYSCAT") + + sys_checks = Table("CHECKS", ischema, + Column("CONSTNAME", CoerceUnicode, key="constname"), + Column("TABSCHEMA", CoerceUnicode, key="tabschema"), + Column("TABNAME", CoerceUnicode, key="tabname"), + Column("TYPE", CoerceUnicode, key="type"), + Column("TEXT", CoerceUnicode, key="text"), schema="SYSCAT") sys_tabconst = Table("TABCONST", ischema, @@ -198,6 +227,7 @@ class DB2Reflector(BaseReflector): Column("TABNAME", CoerceUnicode, key="tabname"), Column("CONSTNAME", CoerceUnicode, key="constname"), Column("COLNAME", CoerceUnicode, key="colname"), + Column("COLSEQ", sa_types.SMALLINT, key="colseq"), schema="SYSCAT") sys_foreignkeys = Table("SQLFOREIGNKEYS", ischema, @@ -227,6 +257,7 @@ class DB2Reflector(BaseReflector): Column("IDENTITY", CoerceUnicode, key="identity"), Column("GENERATED", CoerceUnicode, key="generated"), Column("REMARKS", CoerceUnicode, key="remarks"), + Column("CODEPAGE", sa_types.SMALLINT, key="codepage"), schema="SYSCAT") sys_views = Table("VIEWS", ischema, @@ -276,32 +307,19 @@ def has_table(self, connection, table_name, schema=None, **kw): raise @log_entry_exit - def has_sequence(self, connection, sequence_name, schema=None): - try: - logger.debug(f"Checking sequence existence -> schema={schema}, sequence={sequence_name}") - current_schema = self.denormalize_name(schema or self.default_schema_name) - sequence_name = self.denormalize_name(sequence_name) - logger.debug( - f"Resolved identifiers -> " - f"schema={current_schema}, " - f"sequence={sequence_name}" + def has_sequence(self, connection, sequence_name, schema=None, **kw): + # SQLAlchemy 2's Inspector passes info_cache and other keywords. + current_schema = self.denormalize_name(schema or self.default_schema_name) + sequence_name = self.denormalize_name(sequence_name) + if current_schema: + whereclause = sql.and_( + self.sys_sequences.c.seqschema == current_schema, + self.sys_sequences.c.seqname == sequence_name ) - if current_schema: - whereclause = sql.and_( - self.sys_sequences.c.seqschema == current_schema, - self.sys_sequences.c.seqname == sequence_name - ) - else: - whereclause = self.sys_sequences.c.seqname == sequence_name - s = sql.select(self.sys_sequences.c.seqname).where(whereclause) - logger.debug(f"Generated has_sequence SQL -> {s}") - result = connection.execute(s).first() is not None - logger.debug(f"has_sequence result -> sequence={sequence_name}, exists={result}") - return result - except Exception as e: - logger.error(f"Error checking sequence existence: {e}") - logger.exception("Stack trace in has_sequence") - raise + else: + whereclause = self.sys_sequences.c.seqname == sequence_name + s = sql.select(self.sys_sequences.c.seqname).where(whereclause) + return connection.execute(s).first() is not None @reflection.cache @log_entry_exit @@ -435,106 +453,91 @@ def get_view_definition(self, connection, viewname, schema=None, **kw): @reflection.cache @log_entry_exit def get_columns(self, connection, table_name, schema=None, **kw): - try: - current_schema = self.denormalize_name(schema or self.default_schema_name) - table_name = self.denormalize_name(table_name) - logger.debug(f"Fetching columns -> schema={current_schema}, table={table_name}") - syscols = self.sys_columns - query = ( - sql.select( - syscols.c.colname, syscols.c.typename, - syscols.c.defaultval, syscols.c.nullable, - syscols.c.length, syscols.c.scale, - syscols.c.identity, syscols.c.generated, - syscols.c.remarks - ) - .where(and_( - syscols.c.tabschema == current_schema, - syscols.c.tabname == table_name - )) - .order_by(syscols.c.colno) + current_schema = self.denormalize_name(schema or self.default_schema_name) + table_name = self.denormalize_name(table_name) + syscols = self.sys_columns + query = ( + sql.select( + syscols.c.colname, syscols.c.typename, + syscols.c.defaultval, syscols.c.nullable, + syscols.c.length, syscols.c.scale, + syscols.c.identity, syscols.c.generated, + syscols.c.remarks, syscols.c.codepage ) - logger.debug(f"Generated get_columns SQL -> {query}") - sa_columns = [] - for r in connection.execute(query): - raw_type = r[1].upper() - logger.debug( - f"Processing column -> " - f"name={r[0]}, type={raw_type}, " - f"length={r[4]}, scale={r[5]}" - ) - if raw_type in ['DECIMAL', 'NUMERIC']: - coltype = self.ischema_names.get(raw_type)(int(r[4]), int(r[5])) - elif raw_type in ['CHARACTER', 'CHAR', 'VARCHAR', - 'GRAPHIC', 'VARGRAPHIC']: - coltype = self.ischema_names.get(raw_type)(int(r[4])) - else: - try: - coltype = self.ischema_names[raw_type] - except KeyError: - logger.warning( - f"Unrecognized column type '{raw_type}' " - f"for column '{r[0]}'" - ) - coltype = sa_types.NULLTYPE - column_info = { - 'name': self.normalize_name(r[0]), - 'type': coltype, - 'nullable': r[3] == 'Y', - 'default': r[2] or None, - 'autoincrement': (r[6] == 'Y') and (r[7] != ' '), - 'comment': r[8] or None, - } - logger.debug(f"Column reflected -> {column_info}") - sa_columns.append(column_info) - logger.debug(f"Total columns reflected -> count={len(sa_columns)}") - return sa_columns - except Exception as e: - logger.error(f"Error reflecting columns: {e}") - logger.exception("Stack trace in get_columns") - raise + .where(and_( + syscols.c.tabschema == current_schema, + syscols.c.tabname == table_name + )) + .order_by(syscols.c.colno) + ) + sa_columns = [] + for r in connection.execute(query): + raw_type = r[1].upper() + if raw_type in ('CHARACTER', 'CHAR', 'VARCHAR') and r[9] == 0: + # FOR BIT DATA: the DBAPI returns bytes. + binary = sa_types.VARBINARY if raw_type == 'VARCHAR' else sa_types.BINARY + coltype = binary(int(r[4])) + elif raw_type in ['DECIMAL', 'NUMERIC']: + coltype = self.ischema_names.get(raw_type)(int(r[4]), int(r[5])) + elif raw_type in ['CHARACTER', 'CHAR', 'VARCHAR', + 'GRAPHIC', 'VARGRAPHIC', 'BINARY', 'VARBINARY']: + coltype = self.ischema_names.get(raw_type)(int(r[4])) + elif raw_type == 'DECFLOAT': + # LENGTH is the storage size: 8 bytes for 16 digits, 16 for 34. + coltype = self.ischema_names[raw_type](16 if int(r[4]) == 8 else 34) + else: + try: + coltype = self.ischema_names[raw_type] + except KeyError: + logger.warning( + f"Unrecognized column type '{raw_type}' " + f"for column '{r[0]}'" + ) + coltype = sa_types.NULLTYPE + sa_columns.append({ + 'name': self.normalize_name(r[0]), + 'type': coltype, + 'nullable': r[3] == 'Y', + 'default': r[2] or None, + 'autoincrement': (r[6] == 'Y') and (r[7] != ' '), + 'comment': r[8] or None, + }) + return sa_columns + + def _constraint_columns(self, connection, table_name, schema, kind): + """(constraint name, column name) rows in key order.""" + current_schema = self.denormalize_name(schema or self.default_schema_name) + table_name = self.denormalize_name(table_name) + keycol, const = self.sys_keycoluse, self.sys_tabconst + query = ( + sql.select(keycol.c.constname, keycol.c.colname) + .select_from(join(keycol, const, and_( + keycol.c.tabschema == const.c.tabschema, + keycol.c.tabname == const.c.tabname, + keycol.c.constname == const.c.constname, + ))) + .where(and_( + const.c.tabschema == current_schema, + const.c.tabname == table_name, + const.c.type == kind, + )) + .order_by(keycol.c.constname, keycol.c.colseq) + ) + return [ + (self.normalize_name(name), self.normalize_name(col)) + for name, col in connection.execute(query) + ] @reflection.cache @log_entry_exit def get_pk_constraint(self, connection, table_name, schema=None, **kw): - try: - current_schema = self.denormalize_name(schema or self.default_schema_name) - table_name = self.denormalize_name(table_name) - logger.debug(f"Fetching primary key -> schema={current_schema}, table={table_name}") - sysindexes = self.sys_indexes - col_finder = re.compile(r"(\w+)") - query = ( - sql.select(sysindexes.c.colnames, sysindexes.c.indname) - .where(and_( - sysindexes.c.tabschema == current_schema, - sysindexes.c.tabname == table_name, - sysindexes.c.uniquerule == 'P' - )) - .order_by( - sysindexes.c.tabschema, - sysindexes.c.tabname - )) - logger.debug(f"Generated get_pk_constraint SQL -> {query}") - pk_columns = [] - pk_name = None - for r in connection.execute(query): - cols = col_finder.findall(r[0]) - pk_columns.extend(cols) - if not pk_name: - pk_name = self.normalize_name(r[1]) - normalized_columns = [self.normalize_name(col) for col in pk_columns] - logger.debug( - f"Primary key reflected -> " - f"name={pk_name}, columns={normalized_columns}" - ) - return { - "constrained_columns": normalized_columns, - "name": pk_name - } - except Exception as e: - logger.error(f"Error reflecting primary key: {e}") - logger.exception("Stack trace in get_pk_constraint") - raise + # Read the constraint's own name and columns in key order; splitting + # SYSCAT.INDEXES.COLNAMES on \w+ broke names such as "AMT$X". + rows = self._constraint_columns(connection, table_name, schema, 'P') + return { + "constrained_columns": [col for _, col in rows], + "name": rows[0][0] if rows else None, + } @reflection.cache @log_entry_exit @@ -570,68 +573,64 @@ def get_primary_keys(self, connection, table_name, schema=None, **kw): @reflection.cache @log_entry_exit def get_foreign_keys(self, connection, table_name, schema=None, **kw): - try: - default_schema = self.default_schema_name - current_schema = self.denormalize_name(schema or default_schema) - normalized_default_schema = self.normalize_name(default_schema) - table_name = self.denormalize_name(table_name) - logger.debug( - f"Fetching foreign keys -> " - f"schema={current_schema}, table={table_name}" + # Scope the constrained table by schema and name, and pair columns by + # key position, so a same-named table in another schema does not + # contribute its keys and composite keys keep their order. + default_schema = self.normalize_name(self.default_schema_name) + current_schema = self.denormalize_name(schema or self.default_schema_name) + table_name = self.denormalize_name(table_name) + ref = self.sys_references + fk = self.sys_keycoluse.alias("fk") + pk = self.sys_keycoluse.alias("pk") + query = ( + sql.select( + ref.c.constname, fk.c.colname, ref.c.reftabschema, + ref.c.reftabname, pk.c.colname, ref.c.deleterule, + ref.c.updaterule, ) - sysfkeys = self.sys_foreignkeys - systbl = self.sys_tables - query = ( - sql.select( - sysfkeys.c.fkname, sysfkeys.c.fktabschema, - sysfkeys.c.fktabname, sysfkeys.c.fkcolname, - sysfkeys.c.pkname, sysfkeys.c.pktabschema, - sysfkeys.c.pktabname, sysfkeys.c.pkcolname - ) - .select_from( - join( - systbl, - sysfkeys, - sql.and_( - systbl.c.tabname == sysfkeys.c.pktabname, - systbl.c.tabschema == sysfkeys.c.pktabschema - ) - ) - ) - .where(systbl.c.type == 'T') - .where(systbl.c.tabschema == current_schema) - .where(sysfkeys.c.fktabname == table_name) - .order_by(systbl.c.tabname) + .select_from( + join(ref, fk, and_( + fk.c.tabschema == ref.c.tabschema, + fk.c.tabname == ref.c.tabname, + fk.c.constname == ref.c.constname, + )).join(pk, and_( + pk.c.tabschema == ref.c.reftabschema, + pk.c.tabname == ref.c.reftabname, + pk.c.constname == ref.c.refkeyname, + pk.c.colseq == fk.c.colseq, + )) ) - logger.debug(f"Generated get_foreign_keys SQL -> {query}") - fschema = {} - for r in connection.execute(query): - fk_name = r[0] - if fk_name not in fschema: - referred_schema = self.normalize_name(r[5]) - # if no schema specified and referred schema here is the - # default, then set to None - if schema is None and \ - referred_schema == normalized_default_schema: - referred_schema = None - fschema[fk_name] = { - 'name': self.normalize_name(fk_name), - 'constrained_columns': [self.normalize_name(r[3])], - 'referred_schema': referred_schema, - 'referred_table': self.normalize_name(r[6]), - 'referred_columns': [self.normalize_name(r[7])] - } - logger.debug(f"Foreign key discovered -> {fschema[fk_name]}") - else: - fschema[fk_name]['constrained_columns'].append(self.normalize_name(r[3])) - fschema[fk_name]['referred_columns'].append(self.normalize_name(r[7])) - result = [value for value in fschema.values()] - logger.debug(f"Total foreign keys reflected -> count={len(result)}") - return result - except Exception as e: - logger.error(f"Error reflecting foreign keys: {e}") - logger.exception("Stack trace in get_foreign_keys") - raise + .where(and_( + ref.c.tabschema == current_schema, + ref.c.tabname == table_name, + )) + .order_by(ref.c.constname, fk.c.colseq) + ) + rules = {'C': 'CASCADE', 'N': 'SET NULL', 'R': 'RESTRICT'} + fschema = {} + for name, col, ref_schema, ref_table, ref_col, on_delete, on_update in \ + connection.execute(query): + if name not in fschema: + # SYSCAT.REFERENCES pads schema names. + referred_schema = self.normalize_name(ref_schema.rstrip()) + if schema is None and referred_schema == default_schema: + referred_schema = None + options = {} + if on_delete in rules: + options['ondelete'] = rules[on_delete] + if on_update == 'R': + options['onupdate'] = 'RESTRICT' + fschema[name] = { + 'name': self.normalize_name(name), + 'constrained_columns': [], + 'referred_schema': referred_schema, + 'referred_table': self.normalize_name(ref_table), + 'referred_columns': [], + 'options': options, + } + fschema[name]['constrained_columns'].append(self.normalize_name(col)) + fschema[name]['referred_columns'].append(self.normalize_name(ref_col)) + return list(fschema.values()) @reflection.cache @log_entry_exit @@ -694,126 +693,68 @@ def get_incoming_foreign_keys(self, connection, table_name, schema=None, **kw): @reflection.cache @log_entry_exit def get_indexes(self, connection, table_name, schema=None, **kw): - try: - current_schema = self.denormalize_name(schema or self.default_schema_name) - table_name = self.denormalize_name(table_name) - logger.debug(f"Fetching indexes -> schema={current_schema}, table={table_name}") - sysidx = self.sys_indexes - query = ( - sql.select(sysidx.c.indname, sysidx.c.colnames, - sysidx.c.uniquerule, sysidx.c.system_required - ) - .where(and_( - sysidx.c.tabschema == current_schema, - sysidx.c.tabname == table_name - )) - .order_by(sysidx.c.tabname) - ) - logger.debug(f"Generated get_indexes SQL -> {query}") - indexes = [] - col_finder = re.compile(r"(\w+)") - for r in connection.execute(query): - index_name = r[0] - column_text = r[1] - unique_rule = r[2] - system_required = r[3] - logger.debug( - f"Processing index row -> " - f"name={index_name}, unique_rule={unique_rule}, " - f"system_required={system_required}" - ) - if unique_rule == 'P': - logger.debug(f"Skipping primary key index -> {index_name}") - continue - if unique_rule == 'U' and system_required != 0: - logger.debug(f"Skipping system-required unique index -> {index_name}") - continue - if 'sqlnotapplicable' in column_text.lower(): - logger.debug(f"Skipping internal index -> {index_name}") - continue - normalized_columns = [self.normalize_name(col) for col in col_finder.findall(column_text)] - index_info = { - 'name': self.normalize_name(index_name), - 'column_names': normalized_columns, - 'unique': unique_rule == 'U' - } - logger.debug(f"Index reflected -> {index_info}") - indexes.append(index_info) - logger.debug(f"Total indexes reflected -> count={len(indexes)}") - return indexes - except Exception as e: - logger.error(f"Error reflecting indexes: {e}") - logger.exception("Stack trace in get_indexes") - raise + # Read columns from SYSCAT.INDEXCOLUSE (splitting COLNAMES on \w+ + # broke names such as "AMT$X" and lost DESC), and skip DB2's internal + # XML region/path indexes, which are not user indexes. + current_schema = self.denormalize_name(schema or self.default_schema_name) + table_name = self.denormalize_name(table_name) + idx, cols = self.sys_indexes, self.sys_indexcoluse + query = ( + sql.select(idx.c.indname, idx.c.uniquerule, cols.c.colname, cols.c.colorder) + .select_from(join(idx, cols, and_( + cols.c.indschema == idx.c.indschema, + cols.c.indname == idx.c.indname, + ))) + .where(and_( + idx.c.tabschema == current_schema, + idx.c.tabname == table_name, + idx.c.uniquerule != 'P', + # System-required unique indexes back unique constraints, + # which get_unique_constraints reports. + not_(and_(idx.c.uniquerule == 'U', idx.c.system_required != 0)), + idx.c.indextype.in_(('REG', 'CLUS')), + )) + .order_by(idx.c.indname, cols.c.colseq) + ) + indexes = {} + for name, rule, col, order in connection.execute(query): + index = indexes.setdefault(name, { + 'name': self.normalize_name(name), + 'column_names': [], + 'unique': rule == 'U', + }) + index['column_names'].append(self.normalize_name(col)) + if order == 'D': + index.setdefault('column_sorting', {})[self.normalize_name(col)] = ('desc',) + return list(indexes.values()) @reflection.cache @log_entry_exit def get_unique_constraints(self, connection, table_name, schema=None, **kw): - try: - current_schema = self.denormalize_name(schema or self.default_schema_name) - table_name = self.denormalize_name(table_name) - logger.debug( - f"Fetching unique constraints -> " - f"schema={current_schema}, table={table_name}" - ) - syskeycol = self.sys_keycoluse - sysconst = self.sys_tabconst - query = ( - sql.select( - syskeycol.c.constname, - syskeycol.c.colname - ) - .select_from( - join( - syskeycol, - sysconst, - and_( - syskeycol.c.constname == sysconst.c.constname, - syskeycol.c.tabschema == sysconst.c.tabschema, - syskeycol.c.tabname == sysconst.c.tabname, - ), - ) - ) - .where( - and_( - sysconst.c.tabname == table_name, - sysconst.c.tabschema == current_schema, - sysconst.c.type == "U", - ) - ) - .order_by(syskeycol.c.constname) - ) - logger.debug(f"Generated get_unique_constraints SQL -> {query}") - uniqueConsts = [] - currConst = None - for r in connection.execute(query): - constraint_name = r[0] - column_name = self.normalize_name(r[1]) - if currConst == constraint_name: - uniqueConsts[-1]["column_names"].append(column_name) - logger.debug( - f"Appending column to constraint -> " - f"name={constraint_name}, column={column_name}" - ) - else: - currConst = constraint_name - constraint_info = { - "name": self.normalize_name(currConst), - "column_names": [column_name], - } - logger.debug(f"New unique constraint discovered -> {constraint_info}") - uniqueConsts.append(constraint_info) - logger.debug( - f"Total unique constraints reflected -> " - f"count={len(uniqueConsts)}" - ) - return uniqueConsts - except Exception as e: - logger.error(f"Error reflecting unique constraints: {e}") - logger.exception("Stack trace in get_unique_constraints") - raise - + constraints = {} + for name, col in self._constraint_columns(connection, table_name, schema, 'U'): + constraints.setdefault(name, []).append(col) + return [{'name': k, 'column_names': v} for k, v in constraints.items()] + @reflection.cache + @log_entry_exit + def get_check_constraints(self, connection, table_name, schema=None, **kw): + current_schema = self.denormalize_name(schema or self.default_schema_name) + table_name = self.denormalize_name(table_name) + checks = self.sys_checks + query = ( + sql.select(checks.c.constname, checks.c.text) + .where(and_( + checks.c.tabschema == current_schema, + checks.c.tabname == table_name, + checks.c.type == 'C', + )) + .order_by(checks.c.constname) + ) + return [ + {'name': self.normalize_name(name), 'sqltext': text} + for name, text in connection.execute(query) + ] class AS400Reflector(BaseReflector): ischema = MetaData() diff --git a/test/test_luw_reflection.py b/test/test_luw_reflection.py new file mode 100644 index 0000000..b705052 --- /dev/null +++ b/test/test_luw_reflection.py @@ -0,0 +1,236 @@ +"""DB2 LUW reflection, binary binds, XML and DECFLOAT results with ibm_db.""" + +from decimal import Decimal + +from sqlalchemy import ( + BINARY, + LargeBinary, + MetaData, + Table, + VARBINARY, + inspect, + select, +) +from sqlalchemy.testing import fixtures +from sqlalchemy.testing.assertions import eq_ + +from ibm_db_sa import base +from ibm_db_sa.base import XML +from ibm_db_sa.ibm_db import DB2Dialect_ibm_db +from ibm_db_sa.reflection import DB2Reflector + + +class TestProcessors(fixtures.TestBase): + def test_binary_binds_bytes(self): + dialect = DB2Dialect_ibm_db() + for type_ in (LargeBinary(), BINARY(4), VARBINARY(8)): + processor = type_.dialect_impl(dialect).bind_processor(dialect) + bound = processor(memoryview(b"\x00\xff")) + eq_(bound, b"\x00\xff") + eq_(type(bound), bytes) + eq_(processor(None), None) + + def test_xml_drops_cli_serialization_prefix(self): + dialect = DB2Dialect_ibm_db() + processor = XML().dialect_impl(dialect).result_processor(dialect, None) + prefix = '\ufeff' + eq_(processor(prefix + "1"), "1") + other = '' + eq_(processor(other), other) + eq_(processor(None), None) + + def test_decfloat_processors(self): + dialect = DB2Dialect_ibm_db() + impl = base.DECFLOAT(34).dialect_impl(dialect) + eq_(type(impl), base.DECFLOAT) + value = "3.141592653589793238462643383279502" + result = impl.result_processor(dialect, None)(value) + eq_(result, Decimal(value)) + eq_(type(result), Decimal) + eq_(impl.bind_processor(dialect)(Decimal(value)), value) + eq_(str(base.DECFLOAT(16).compile(dialect=dialect)), "DECFLOAT(16)") + + def test_is_disconnect_covers_base_error(self): + dialect = DB2Dialect_ibm_db() + dbapi = dialect.dbapi = DB2Dialect_ibm_db.import_dbapi() + lost = "[IBM][CLI Driver] SQL30081N A communication error has been detected." + eq_(dialect.is_disconnect(dbapi.Error(lost), None, None), True) + eq_(dialect.is_disconnect(dbapi.Error("SQL0204N"), None, None), False) + + def test_has_sequence_accepts_inspector_keywords(self): + dialect = DB2Dialect_ibm_db() + reflector = DB2Reflector(dialect) + + class Result: + def first(self): + return ("SEQ1",) + + class Connection: + def execute(self, statement): + return Result() + + eq_( + reflector.has_sequence(Connection(), "seq1", schema="s", info_cache={}), + True, + ) + + +class TestLUWReflection(fixtures.TestBase): + __only_on__ = "ibm_db_sa+ibm_db_sa" + __backend__ = True + + DDL = [ + "CREATE SCHEMA IBMSA_A", + "CREATE SCHEMA IBMSA_B", + "CREATE TABLE IBMSA_A.PARENT (A INT NOT NULL, B INT NOT NULL, " + "CONSTRAINT PK_PARENT PRIMARY KEY (B, A))", + 'CREATE TABLE IBMSA_A.CHILD (ID INT NOT NULL, PA INT, PB INT, "AMT$X" INT ' + "NOT NULL, U1 INT NOT NULL, U2 INT NOT NULL, X XML, " + 'CONSTRAINT PK_CHILD PRIMARY KEY (ID, "AMT$X"), ' + "CONSTRAINT FK_PARENT FOREIGN KEY (PB, PA) REFERENCES IBMSA_A.PARENT (B, A) " + "ON DELETE CASCADE, CONSTRAINT UQ_C UNIQUE (U2, U1), " + "CONSTRAINT CK_U CHECK (U1 >= 0))", + "CREATE TABLE IBMSA_B.PARENT (A INT NOT NULL PRIMARY KEY)", + "CREATE TABLE IBMSA_B.CHILD (ID INT NOT NULL PRIMARY KEY, PA INT NOT NULL, " + "CONSTRAINT FK_OTHER FOREIGN KEY (PA) REFERENCES IBMSA_B.PARENT (A), " + "CONSTRAINT UQ_C UNIQUE (PA))", + 'CREATE INDEX IBMSA_A.IX_CHILD ON IBMSA_A.CHILD (U1 DESC, "AMT$X" ASC)', + "CREATE SEQUENCE IBMSA_A.SEQ1", + "CREATE TABLE IBMSA_A.TYPES (ID INT NOT NULL, DF DECFLOAT(16), " + "BI BINARY(4), VB VARBINARY(8), CB CHAR(4) FOR BIT DATA, BL BLOB(1K))", + ] + DROP = [ + "DROP TABLE IBMSA_A.CHILD", + "DROP TABLE IBMSA_A.PARENT", + "DROP TABLE IBMSA_A.TYPES", + "DROP TABLE IBMSA_B.CHILD", + "DROP TABLE IBMSA_B.PARENT", + "DROP SEQUENCE IBMSA_A.SEQ1", + "DROP SCHEMA IBMSA_A RESTRICT", + "DROP SCHEMA IBMSA_B RESTRICT", + ] + + @classmethod + def setup_test_class(cls): + from sqlalchemy.testing import config + + with config.db.begin() as conn: + for statement in cls.DDL: + conn.exec_driver_sql(statement) + + @classmethod + def teardown_test_class(cls): + from sqlalchemy.testing import config + + for statement in cls.DROP: + with config.db.begin() as conn: + conn.exec_driver_sql(statement) + + def test_foreign_keys_are_schema_scoped_and_ordered(self, connection): + inspector = inspect(connection) + eq_( + inspector.get_foreign_keys("child", schema="ibmsa_a"), + [ + { + "name": "fk_parent", + "constrained_columns": ["pb", "pa"], + "referred_schema": "ibmsa_a", + "referred_table": "parent", + "referred_columns": ["b", "a"], + "options": {"ondelete": "CASCADE"}, + } + ], + ) + eq_( + inspector.get_foreign_keys("child", schema="ibmsa_b"), + [ + { + "name": "fk_other", + "constrained_columns": ["pa"], + "referred_schema": "ibmsa_b", + "referred_table": "parent", + "referred_columns": ["a"], + "options": {}, + } + ], + ) + + def test_unique_constraints_are_table_scoped_and_ordered(self, connection): + inspector = inspect(connection) + eq_( + inspector.get_unique_constraints("child", schema="ibmsa_a"), + [{"name": "uq_c", "column_names": ["u2", "u1"]}], + ) + eq_( + inspector.get_unique_constraints("child", schema="ibmsa_b"), + [{"name": "uq_c", "column_names": ["pa"]}], + ) + + def test_primary_key_constraint_name_and_columns(self, connection): + eq_( + inspect(connection).get_pk_constraint("child", schema="ibmsa_a"), + {"constrained_columns": ["id", "amt$x"], "name": "pk_child"}, + ) + + def test_indexes(self, connection): + eq_( + inspect(connection).get_indexes("child", schema="ibmsa_a"), + [ + { + "name": "ix_child", + "column_names": ["u1", "amt$x"], + "unique": False, + "column_sorting": {"u1": ("desc",)}, + } + ], + ) + + def test_check_constraints(self, connection): + eq_( + inspect(connection).get_check_constraints("child", schema="ibmsa_a"), + [{"name": "ck_u", "sqltext": "U1 >= 0"}], + ) + + def test_has_sequence(self, connection): + inspector = inspect(connection) + eq_(inspector.has_sequence("seq1", schema="ibmsa_a"), True) + eq_(inspector.has_sequence("nope", schema="ibmsa_a"), False) + + def test_column_types_match_values(self, connection): + columns = { + c["name"]: c["type"] + for c in inspect(connection).get_columns("types", schema="ibmsa_a") + } + eq_(repr(columns["df"]), "DECFLOAT(precision=16)") + eq_(repr(columns["bi"]), "BINARY(length=4)") + eq_(repr(columns["vb"]), "VARBINARY(length=8)") + eq_(repr(columns["cb"]), "BINARY(length=4)") + table = Table("types", MetaData(), schema="ibmsa_a", autoload_with=connection) + rows = [ + dict( + id=1, + df=Decimal("1.234567890123456"), + bi=b"abcd", + vb=b"\x01", + cb=b"\xde\xad\xbe\xef", + bl=b"\x00\xff", + ), + dict( + id=2, df=Decimal("-0.5"), bi=b"wxyz", vb=b"", cb=b"\x00" * 4, bl=b"\x10" + ), + dict(id=3, df=None, bi=None, vb=None, cb=None, bl=None), + ] + connection.execute(table.insert(), rows) + actual = [ + dict(r._mapping) + for r in connection.execute(select(table).order_by(table.c.id)) + ] + eq_(actual, rows) + eq_([type(r["df"]) for r in actual], [Decimal, Decimal, type(None)]) + + def test_xml_value(self, connection): + table = Table("child", MetaData(), schema="ibmsa_a", autoload_with=connection) + connection.execute( + table.insert(), [{"id": 1, "amt$x": 0, "u1": 1, "u2": 1, "x": "1"}] + ) + eq_(connection.execute(select(table.c.x)).scalar(), "1")