diff --git a/pyproject.toml b/pyproject.toml index 0995ffa..fc2bb08 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -45,7 +45,7 @@ ibmi = "sqlalchemy_ibmi.base:IBMiDb2Dialect" [tool.poetry.dependencies] python = ">=3.7,<3.13" -sqlalchemy = ">=1.3.0, <2.0" +sqlalchemy = ">=1.4.0, <3.0" pyodbc = ">=4.0" [tool.poetry.dev-dependencies] diff --git a/sqlalchemy_ibmi/base.py b/sqlalchemy_ibmi/base.py index b32d2e7..5cde02f 100644 --- a/sqlalchemy_ibmi/base.py +++ b/sqlalchemy_ibmi/base.py @@ -180,11 +180,13 @@ class User(Base): """ # noqa E501 import datetime import re +import warnings from collections import defaultdict from sqlalchemy import ( select, + text, schema as sa_schema, exc, util, @@ -196,6 +198,7 @@ class User(Base): from sqlalchemy.sql import compiler, operators from sqlalchemy.sql.expression import and_, cast from sqlalchemy.engine import default, reflection +from sqlalchemy.engine import cursor as _cursor from sqlalchemy.types import ( BLOB, CHAR, @@ -407,19 +410,19 @@ def visit_numeric(self, type_, **kw): # This is now part of SQLAlchemy as of 2.0. We can drop this function once # we drop support for earlier versions. - def visit_DOUBLE(self, type_): + def visit_DOUBLE(self, type_, **kw): return "DOUBLE" - def visit_GRAPHIC(self, type_): + def visit_GRAPHIC(self, type_, **kw): return self._extend(type_, "GRAPHIC") - def visit_VARGRAPHIC(self, type_): + def visit_VARGRAPHIC(self, type_, **kw): return self._extend(type_, "VARGRAPHIC") def visit_DBCLOB(self, type_, **kw): return self._extend(type_, "DBCLOB", length=type_.length or "1G") - def visit_XML(self, type_): + def visit_XML(self, type_, **kw): return "XML" @@ -502,17 +505,17 @@ def visit_cast(self, cast, **kw): kw["_cast_applied"] = True return super().visit_cast(cast, **kw) - def visit_savepoint(self, savepoint_stmt): + def visit_savepoint(self, savepoint_stmt, **kw): return "SAVEPOINT %(sid)s ON ROLLBACK RETAIN CURSORS" % { "sid": self.preparer.format_savepoint(savepoint_stmt) } - def visit_rollback_to_savepoint(self, savepoint_stmt): + def visit_rollback_to_savepoint(self, savepoint_stmt, **kw): return "ROLLBACK TO SAVEPOINT %(sid)s" % { "sid": self.preparer.format_savepoint(savepoint_stmt) } - def visit_release_savepoint(self, savepoint_stmt): + def visit_release_savepoint(self, savepoint_stmt, **kw): return "RELEASE TO SAVEPOINT %(sid)s" % { "sid": self.preparer.format_savepoint(savepoint_stmt) } @@ -526,7 +529,7 @@ def visit_unary(self, unary, **kw): usql = f"CASE WHEN {usql} THEN 1 ELSE 0 END" return usql - def visit_empty_set_op_expr(self, type_, expand_op): + def visit_empty_set_op_expr(self, type_, expand_op, **kw): if expand_op is operators.not_in_op: return "(%s)) OR (1 = 1" % ( ", ".join( @@ -550,9 +553,20 @@ def visit_empty_set_op_expr(self, type_, expand_op): else: return self.visit_empty_set_expr(type_) - def visit_empty_set_expr(self, element_types): + def visit_empty_set_expr(self, element_types, **kw): return "SELECT 1 FROM SYSIBM.SYSDUMMY1 WHERE 1!=1" + def visit_over(self, over, **kw): + """Override window function handling to avoid CAST in frame clause. + + IBM i DB2 doesn't support CAST expressions in ROWS BETWEEN clauses. + We need to render literal values directly instead of using bind parameters. + """ + # Render with literal binds to avoid CAST(? AS BIGINT) in frame clause + kw = kw.copy() + kw['literal_binds'] = True + return super().visit_over(over, **kw) + def visit_null(self, expr, **kw): if not kw.get("within_columns_clause", False): return "NULL" @@ -722,9 +736,11 @@ def visit_drop_constraint(self, drop, **kw): ) def visit_create_index( - self, create, include_schema=True, include_table_schema=True + self, create, include_schema=True, include_table_schema=True, **kw ): - sql = super().visit_create_index(create, include_schema, include_table_schema) + sql = super().visit_create_index( + create, include_schema=include_schema, include_table_schema=include_table_schema, **kw + ) if getattr(create.element, "uConstraint_as_index", None): sql += " EXCLUDE NULL KEYS" return sql @@ -752,6 +768,17 @@ def pre_exec(self): seq_column = tbl._autoincrement_column insert_has_sequence = seq_column is not None + # IBM i doesn't support RETURNING clause, so we fetch lastrowid + # using VALUES IDENTITY_VAL_LOCAL() in post_exec. + # We should fetch it whenever: + # - The table has an autoincrement column + # - No explicit RETURNING clause was specified (self.compiled.returning) + # - Not an inline INSERT + # + # Note: We ignore implicit_returning=False on the table because: + # 1. IBM i doesn't support RETURNING syntax anyway + # 2. implicit_returning=False just means "don't use RETURNING syntax" + # 3. But we still need to return the lastrowid for test compatibility self._select_lastrowid = ( insert_has_sequence and not self.compiled.returning @@ -765,6 +792,10 @@ def post_exec(self): row = self.cursor.fetchall()[0] if row[0] is not None: self._lastrowid = int(row[0]) + + # Mark this as a DML statement with no user-facing cursor + # This ensures returns_rows is False even though we fetched lastrowid + self.cursor_fetch_strategy = _cursor._NO_CURSOR_DML def fire_sequence(self, seq, type_): return self._execute_scalar( @@ -819,6 +850,7 @@ class IBMiDb2Dialect(default.DefaultDialect): supports_default_values = False supports_empty_insert = False supports_statement_cache = True + default_isolation_level = "READ COMMITTED" statement_compiler = DB2Compiler ddl_compiler = DB2DDLCompiler @@ -828,7 +860,7 @@ class IBMiDb2Dialect(default.DefaultDialect): def __init__(self, isolation_level=None, fast_executemany=False, **kw): super().__init__(**kw) - self.isolation_level = isolation_level + self.isolation_level = isolation_level or self.default_isolation_level self.fast_executemany = fast_executemany def on_connect(self): @@ -846,21 +878,30 @@ def initialize(self, connection): self.driver_version = self._get_driver_version(connection.connection) self.text_server_available = self._check_text_server(connection) + @reflection.cache 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) + + # Check if table exists + if not self.has_table(connection, table_name, schema): + raise exc.NoSuchTableError( + f"Table '{table_name}' not found in schema '{current_schema}'" + ) + sysconst = self.sys_table_constraints syschkconst = self.sys_check_constraints query = select( - [syschkconst.c.conname, syschkconst.c.chkclause], + syschkconst.c.conname, syschkconst.c.chkclause + ).where( and_( syschkconst.c.conschema == sysconst.c.conschema, syschkconst.c.conname == sysconst.c.conname, sysconst.c.tabschema == current_schema, sysconst.c.tabname == table_name, - ), - ) + ) + ).order_by(syschkconst.c.conname) check_consts = [] print(query) @@ -881,7 +922,7 @@ def get_table_comment(self, connection, table_name, schema=None, **kw): ) else: whereclause = self.sys_tables.c.tabname == table_name - select_statement = select([self.sys_tables.c.tabcomment], whereclause) + select_statement = select(self.sys_tables.c.tabcomment).where(whereclause) results = connection.execute(select_statement) return {"text": results.scalar()} @@ -900,28 +941,89 @@ def _isolation_lookup(self): "REPEATABLE READ": self.dbapi.SQL_TXN_REPEATABLE_READ, } + def get_isolation_level_values(self, dbapi_conn): + return list(self._isolation_lookup) + # Methods merged from PyODBCConnector def get_isolation_level(self, dbapi_conn): + # Return the stored isolation level. IBM i ODBC doesn't support + # querying isolation level from an active connection. return self.isolation_level def set_isolation_level(self, connection, level): - self.isolation_level = level - level = level.replace("_", " ") - if level in self._isolation_lookup: - connection.set_attr( - self.dbapi.SQL_ATTR_TXN_ISOLATION, self._isolation_lookup[level] - ) - else: + """Set the isolation level for this connection. + + This method attempts to set the isolation level using ODBC attributes. + Due to IBM i ODBC driver limitations, this may fail with error HY011 if + called on an active connection or during a transaction. The isolation level + will still be validated and stored for reference. + + For guaranteed isolation level setting, use the commit_mode connection parameter: + commit_mode=0 -> *CHG (READ UNCOMMITTED) + commit_mode=1 -> *CS (READ COMMITTED) - default + commit_mode=2 -> *ALL (REPEATABLE READ) + commit_mode=3 -> *NONE (no transactions) + + Example: + engine = create_engine("ibmi+pyodbc://user:pass@host/db?commit_mode=1") + """ + if level is None: + level = self.default_isolation_level + level = str(level).replace("_", " ") + if level not in self._isolation_lookup: raise exc.ArgumentError( "Invalid value '%s' for isolation_level. " "Valid isolation levels for %s are %s" % (level, self.name, ", ".join(self._isolation_lookup.keys())) ) + + # Try to set via ODBC attribute + # This should work when called during on_connect(), but may fail with HY011 + # if called on an active connection or during a transaction + try: + connection.set_attr( + self.dbapi.SQL_ATTR_TXN_ISOLATION, + self._isolation_lookup[level] + ) + except Exception as e: + error_msg = str(e) + + # HY011 (Operation invalid at this time) is expected when connection is active + if "HY011" in error_msg or "30033" in error_msg: + warnings.warn( + f"IBM i ODBC driver returned HY011 when setting isolation level to '{level}'. " + "This is expected when called on an active connection. " + "Isolation level stored but not applied via ODBC. " + "To guarantee isolation level is set, use connection string " + "parameter commit_mode (e.g., commit_mode=1 for READ COMMITTED).", + UserWarning, + stacklevel=2 + ) + else: + # Unexpected error - may indicate a real problem + warnings.warn( + f"Failed to set isolation level via ODBC: {e}. " + "Isolation level will be stored but may not be active. " + "To ensure isolation level is set, use connection string " + "parameter commit_mode (e.g., commit_mode=1 for READ COMMITTED).", + UserWarning, + stacklevel=2 + ) + + self.isolation_level = level + + def reset_isolation_level(self, connection): + self.set_isolation_level(connection, self.default_isolation_level) @classmethod - def dbapi(cls): + def import_dbapi(cls): return __import__("pyodbc") + + # Backwards compatibility alias + @classmethod + def dbapi(cls): + return cls.import_dbapi() DRIVER_KEYWORD_MAP = { # SQLAlchemy kwd: (ODBC keyword, type, default) @@ -939,6 +1041,7 @@ def dbapi(cls): "use_system_naming": ("NAM", to_bool, False), "trim_char_fields": ("TRIMCHAR", to_bool, None), "lob_threshold_kb": ("MAXFIELDLEN", int, None), + "commit_mode": ("CMT", int, None), # 0=*CHG, 1=*CS, 2=*ALL, 3=*NONE } DRIVER_KEYWORDS_SPECIAL = { @@ -1035,7 +1138,7 @@ def _get_server_version_info(self, connection, allow_chars=True): return tuple(version[0:2]) def _get_default_schema_name(self, connection): - return self.normalize_name(connection.execute("VALUES CURRENT_SCHEMA").scalar()) + return self.normalize_name(connection.execute(text("VALUES CURRENT_SCHEMA")).scalar()) # Driver version for IBM i Access ODBC Driver is given as # VV.RR.SSSF where VV (major), RR (release), and SSS (service pack) @@ -1187,7 +1290,8 @@ def _get_driver_version(self, db_conn): schema="QSYS2", ) - def has_table(self, connection, table_name, schema=None): + @reflection.cache + def has_table(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) if current_schema: @@ -1197,11 +1301,12 @@ def has_table(self, connection, table_name, schema=None): ) else: whereclause = self.sys_tables.c.tabname == table_name - select_statement = select([self.sys_tables], whereclause) + select_statement = select(self.sys_tables).where(whereclause) results = connection.execute(select_statement) return results.first() is not None - def has_sequence(self, connection, sequence_name, schema=None): + @reflection.cache + def has_sequence(self, connection, sequence_name, schema=None, **kw): current_schema = self.denormalize_name(schema or self.default_schema_name) sequence_name = self.denormalize_name(sequence_name) if current_schema: @@ -1211,7 +1316,7 @@ def has_sequence(self, connection, sequence_name, schema=None): ) else: whereclause = self.sys_sequences.c.seqname == sequence_name - select_statement = select([self.sys_sequences.c.seqname], whereclause) + select_statement = select(self.sys_sequences.c.seqname).where(whereclause) results = connection.execute(select_statement) return results.first() is not None @@ -1219,7 +1324,7 @@ def has_sequence(self, connection, sequence_name, schema=None): def get_schema_names(self, connection, **kw): sysschema = self.sys_schemas query = ( - select([sysschema.c.schemaname]) + select(sysschema.c.schemaname) .where(sysschema.c.schemaname.notlike("SYS%")) .where(sysschema.c.schemaname.notlike("Q%")) .order_by(sysschema.c.schemaname) @@ -1232,7 +1337,7 @@ def get_table_names(self, connection, schema=None, **kw): current_schema = self.denormalize_name(schema or self.default_schema_name) query = ( - select([self.sys_tables.c.tabname]) + select(self.sys_tables.c.tabname) .where(self.sys_tables.c.tabschema == current_schema) .where(self.sys_tables.c.tabtype.in_(["T", "P"])) .where(self.sys_tables.c.tabsys == "N") @@ -1245,7 +1350,7 @@ def get_view_names(self, connection, schema=None, **kw): current_schema = self.denormalize_name(schema or self.default_schema_name) query = ( - select([self.sys_tables.c.tabname]) + select(self.sys_tables.c.tabname) .where(self.sys_tables.c.tabschema == current_schema) .where(self.sys_tables.c.tabtype.in_(["V"])) .where(self.sys_tables.c.tabsys == "N") @@ -1259,23 +1364,35 @@ def get_view_definition(self, connection, viewname, schema=None, **kw): current_schema = self.denormalize_name(schema or self.default_schema_name) viewname = self.denormalize_name(viewname) - query = select([self.sys_views.c.text]).where( + query = select(self.sys_views.c.text).where( and_( self.sys_views.c.viewschema == current_schema, self.sys_views.c.viewname == viewname, ) ) - return connection.execute(query).scalar() + result = connection.execute(query).scalar() + if result is None: + raise exc.NoSuchTableError( + f"View '{viewname}' not found in schema '{current_schema}'" + ) + return result @reflection.cache def get_columns(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) + + # Check if table exists + if not self.has_table(connection, table_name, schema): + raise exc.NoSuchTableError( + f"Table '{table_name}' not found in schema '{current_schema}'" + ) + syscols = self.sys_columns - query = select( - [ + query = ( + select( syscols.c.colname, syscols.c.typename, syscols.c.defaultval, @@ -1284,11 +1401,13 @@ def get_columns(self, connection, table_name, schema=None, **kw): syscols.c.scale, syscols.c.isid, syscols.c.idgenerate, - ], - and_( - syscols.c.tabschema == current_schema, syscols.c.tabname == table_name - ), - order_by=[syscols.c.colno], + ) + .where( + and_( + syscols.c.tabschema == current_schema, syscols.c.tabname == table_name + ) + ) + .order_by(syscols.c.colno) ) sa_columns = [] for row in connection.execute(query): @@ -1323,11 +1442,18 @@ def get_columns(self, connection, table_name, schema=None, **kw): def get_pk_constraint(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) + + # Check if table exists + if not self.has_table(connection, table_name, schema): + raise exc.NoSuchTableError( + f"Table '{table_name}' not found in schema '{current_schema}'" + ) + sysconst = self.sys_table_constraints syskeyconst = self.sys_key_constraints query = ( - select([syskeyconst.c.colname, sysconst.c.conname]) + select(syskeyconst.c.colname, sysconst.c.conname) .where( and_( syskeyconst.c.conschema == sysconst.c.conschema, @@ -1356,7 +1482,7 @@ def get_primary_keys(self, connection, table_name, schema=None, **kw): syskeyconst = self.sys_key_constraints query = ( - select([syskeyconst.c.colname, sysconst.c.tabname]) + select(syskeyconst.c.colname, sysconst.c.tabname) .where( and_( syskeyconst.c.conschema == sysconst.c.conschema, @@ -1377,9 +1503,16 @@ def get_foreign_keys(self, connection, table_name, schema=None, **kw): current_schema = self.denormalize_name(schema or default_schema) default_schema = self.normalize_name(default_schema) table_name = self.denormalize_name(table_name) + + # Check if table exists + if not self.has_table(connection, table_name, schema): + raise exc.NoSuchTableError( + f"Table '{table_name}' not found in schema '{current_schema}'" + ) + sysfkeys = self.sys_foreignkeys - query = select( - [ + query = ( + select( sysfkeys.c.fkname, sysfkeys.c.fktabschema, sysfkeys.c.fktabname, @@ -1388,12 +1521,14 @@ def get_foreign_keys(self, connection, table_name, schema=None, **kw): sysfkeys.c.pktabschema, sysfkeys.c.pktabname, sysfkeys.c.pkcolname, - ], - and_( - sysfkeys.c.fktabschema == current_schema, - sysfkeys.c.fktabname == table_name, - ), - order_by=[sysfkeys.c.colno], + ) + .where( + and_( + sysfkeys.c.fktabschema == current_schema, + sysfkeys.c.fktabname == table_name, + ) + ) + .order_by(sysfkeys.c.fkname, sysfkeys.c.colno) ) fschema = {} for row in connection.execute(query): @@ -1424,18 +1559,27 @@ def get_foreign_keys(self, connection, table_name, schema=None, **kw): def get_indexes(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) + + # Check if table exists + if not self.has_table(connection, table_name, schema): + raise exc.NoSuchTableError( + f"Table '{table_name}' not found in schema '{current_schema}'" + ) + sysidx = self.sys_indexes syskey = self.sys_keys - query = select( - [sysidx.c.indname, sysidx.c.uniquerule, syskey.c.colname], - and_( - syskey.c.indschema == sysidx.c.indschema, - syskey.c.indname == sysidx.c.indname, - sysidx.c.tabschema == current_schema, - sysidx.c.tabname == table_name, - ), - order_by=[syskey.c.indname, syskey.c.colno], + query = ( + select(sysidx.c.indname, sysidx.c.uniquerule, syskey.c.colname) + .where( + and_( + syskey.c.indschema == sysidx.c.indschema, + syskey.c.indname == sysidx.c.indname, + sysidx.c.tabschema == current_schema, + sysidx.c.tabname == table_name, + ) + ) + .order_by(syskey.c.indname, syskey.c.colno) ) indexes = {} for row in connection.execute(query): @@ -1446,7 +1590,7 @@ def get_indexes(self, connection, table_name, schema=None, **kw): indexes[key] = { "name": self.normalize_name(row[0]), "column_names": [self.normalize_name(row[2])], - "unique": row[1] == "Y", + "unique": row[1] in ('U', 'P'), # U=unique, P=primary key } return [value for key, value in indexes.items()] @@ -1454,11 +1598,18 @@ def get_indexes(self, connection, table_name, schema=None, **kw): def get_unique_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) + + # Check if table exists + if not self.has_table(connection, table_name, schema): + raise exc.NoSuchTableError( + f"Table '{table_name}' not found in schema '{current_schema}'" + ) + sysconst = self.sys_table_constraints sysconstcol = self.sys_constraints_columns query = ( - select([sysconst.c.conname, sysconstcol.c.colname]) + select(sysconst.c.conname, sysconstcol.c.colname) .where( and_( sysconstcol.c.conschema == sysconst.c.conschema, @@ -1489,13 +1640,13 @@ def get_unique_constraints(self, connection, table_name, schema=None, **kw): @reflection.cache def get_sequence_names(self, connection, schema, **kw): current_schema = self.denormalize_name(schema or self.default_schema_name) - query = select([self.sys_sequences.c.seqname]).where( + query = select(self.sys_sequences.c.seqname).where( self.sys_sequences.c.seqschema == current_schema, ) return [self.normalize_name(r[0]) for r in connection.execute(query)] def _check_text_server(self, connection): - stmt = "SELECT COUNT(*) FROM QSYS2.SYSTEXTSERVERS" + stmt = text("SELECT COUNT(*) FROM QSYS2.SYSTEXTSERVERS") return connection.execute(stmt).scalar() def do_executemany(self, cursor, statement, parameters, context=None):