diff options
| author | Mike Bayer <mike_mp@zzzcomputing.com> | 2018-05-23 16:22:48 -0400 |
|---|---|---|
| committer | Mike Bayer <mike_mp@zzzcomputing.com> | 2018-05-25 10:29:10 -0400 |
| commit | 6c27bf5048fc7335a1d1fdd49b651b1a164b8e32 (patch) | |
| tree | 3b447956383f636fa42d343fe5d816ac196ef5bc /lib/sqlalchemy | |
| parent | c7ae04d1c5c4aa6c6099584ae386d6ab9ef7b290 (diff) | |
| download | sqlalchemy-6c27bf5048fc7335a1d1fdd49b651b1a164b8e32.tar.gz | |
Turn oracle BINARY_DOUBLE, BINARY_FLOAT, DOUBLE_PRECISION into floats
The Oracle BINARY_FLOAT and BINARY_DOUBLE datatypes now participate within
cx_Oracle.setinputsizes(), passing along NATIVE_FLOAT, so as to support the
NaN value. Additionally, :class:`.oracle.BINARY_FLOAT`,
:class:`.oracle.BINARY_DOUBLE` and :class:`.oracle.DOUBLE_PRECISION` now
subclass :class:`.Float`, since these are floating point datatypes, not
decimal. These datatypes were already defaulting the
:paramref:`.Float.asdecimal` flag to False in line with what
:class:`.Float` already does.
Added reflection capabilities for the :class:`.oracle.BINARY_FLOAT`,
:class:`.oracle.BINARY_DOUBLE` datatypes.
Change-Id: Id99b912e83052654a17d07dc92b4dcb958cb7600
Fixes: #4264
Diffstat (limited to 'lib/sqlalchemy')
| -rw-r--r-- | lib/sqlalchemy/dialects/oracle/base.py | 44 | ||||
| -rw-r--r-- | lib/sqlalchemy/dialects/oracle/cx_oracle.py | 18 | ||||
| -rw-r--r-- | lib/sqlalchemy/engine/default.py | 2 | ||||
| -rw-r--r-- | lib/sqlalchemy/sql/type_api.py | 30 |
4 files changed, 64 insertions, 30 deletions
diff --git a/lib/sqlalchemy/dialects/oracle/base.py b/lib/sqlalchemy/dialects/oracle/base.py index e55a9cbc6..39acbf28d 100644 --- a/lib/sqlalchemy/dialects/oracle/base.py +++ b/lib/sqlalchemy/dialects/oracle/base.py @@ -411,38 +411,17 @@ class NUMBER(sqltypes.Numeric, sqltypes.Integer): return sqltypes.Integer -class DOUBLE_PRECISION(sqltypes.Numeric): +class DOUBLE_PRECISION(sqltypes.Float): __visit_name__ = 'DOUBLE_PRECISION' - def __init__(self, precision=None, scale=None, asdecimal=None): - if asdecimal is None: - asdecimal = False - - super(DOUBLE_PRECISION, self).__init__( - precision=precision, scale=scale, asdecimal=asdecimal) - -class BINARY_DOUBLE(sqltypes.Numeric): +class BINARY_DOUBLE(sqltypes.Float): __visit_name__ = 'BINARY_DOUBLE' - def __init__(self, precision=None, scale=None, asdecimal=None): - if asdecimal is None: - asdecimal = False - - super(BINARY_DOUBLE, self).__init__( - precision=precision, scale=scale, asdecimal=asdecimal) - -class BINARY_FLOAT(sqltypes.Numeric): +class BINARY_FLOAT(sqltypes.Float): __visit_name__ = 'BINARY_FLOAT' - def __init__(self, precision=None, scale=None, asdecimal=None): - if asdecimal is None: - asdecimal = False - - super(BINARY_FLOAT, self).__init__( - precision=precision, scale=scale, asdecimal=asdecimal) - class BFILE(sqltypes.LargeBinary): __visit_name__ = 'BFILE' @@ -536,6 +515,8 @@ ischema_names = { 'FLOAT': FLOAT, 'DOUBLE PRECISION': DOUBLE_PRECISION, 'LONG': LONG, + 'BINARY_DOUBLE': BINARY_DOUBLE, + 'BINARY_FLOAT': BINARY_FLOAT } @@ -585,17 +566,25 @@ class OracleTypeCompiler(compiler.GenericTypeCompiler): def visit_BINARY_FLOAT(self, type_, **kw): return self._generate_numeric(type_, "BINARY_FLOAT", **kw) + def visit_FLOAT(self, type_, **kw): + # don't support conversion between decimal/binary + # precision yet + kw['no_precision'] = True + return self._generate_numeric(type_, "FLOAT", **kw) + def visit_NUMBER(self, type_, **kw): return self._generate_numeric(type_, "NUMBER", **kw) - def _generate_numeric(self, type_, name, precision=None, scale=None, **kw): + def _generate_numeric( + self, type_, name, precision=None, + scale=None, no_precision=False, **kw): if precision is None: precision = type_.precision if scale is None: scale = getattr(type_, 'scale', None) - if precision is None: + if no_precision or precision is None: return name elif scale is None: n = "%(name)s(%(precision)s)" @@ -1418,6 +1407,9 @@ class OracleDialect(default.DefaultDialect): coltype = INTEGER() else: coltype = NUMBER(precision, scale) + elif coltype == 'FLOAT': + # TODO: support "precision" here as "binary_precision" + coltype = FLOAT() elif coltype in ('VARCHAR2', 'NVARCHAR2', 'CHAR'): coltype = self.ischema_names.get(coltype)(length) elif 'WITH TIME ZONE' in coltype: diff --git a/lib/sqlalchemy/dialects/oracle/cx_oracle.py b/lib/sqlalchemy/dialects/oracle/cx_oracle.py index 0bd682d19..2fbb2074c 100644 --- a/lib/sqlalchemy/dialects/oracle/cx_oracle.py +++ b/lib/sqlalchemy/dialects/oracle/cx_oracle.py @@ -240,7 +240,6 @@ class _OracleInteger(sqltypes.Integer): return handler - class _OracleNumeric(sqltypes.Numeric): is_number = False @@ -323,6 +322,19 @@ class _OracleNumeric(sqltypes.Numeric): return handler +class _OracleBinaryFloat(_OracleNumeric): + def get_dbapi_type(self, dbapi): + return dbapi.NATIVE_FLOAT + + +class _OracleBINARY_FLOAT(_OracleBinaryFloat, oracle.BINARY_FLOAT): + pass + + +class _OracleBINARY_DOUBLE(_OracleBinaryFloat, oracle.BINARY_DOUBLE): + pass + + class _OracleNUMBER(_OracleNumeric): is_number = True @@ -597,6 +609,8 @@ class OracleDialect_cx_oracle(OracleDialect): colspecs = { sqltypes.Numeric: _OracleNumeric, sqltypes.Float: _OracleNumeric, + oracle.BINARY_FLOAT: _OracleBINARY_FLOAT, + oracle.BINARY_DOUBLE: _OracleBINARY_DOUBLE, sqltypes.Integer: _OracleInteger, oracle.NUMBER: _OracleNUMBER, @@ -654,7 +668,7 @@ class OracleDialect_cx_oracle(OracleDialect): cx_Oracle.NCLOB, cx_Oracle.CLOB, cx_Oracle.LOB, cx_Oracle.NCHAR, cx_Oracle.FIXED_NCHAR, cx_Oracle.BLOB, cx_Oracle.FIXED_CHAR, cx_Oracle.TIMESTAMP, - _OracleInteger + _OracleInteger, _OracleBINARY_FLOAT, _OracleBINARY_DOUBLE } self._is_cx_oracle_6 = self.cx_oracle_ver >= (6, ) diff --git a/lib/sqlalchemy/engine/default.py b/lib/sqlalchemy/engine/default.py index ea806deae..4d5f338bf 100644 --- a/lib/sqlalchemy/engine/default.py +++ b/lib/sqlalchemy/engine/default.py @@ -1129,7 +1129,7 @@ class DefaultExecutionContext(interfaces.ExecutionContext): key_to_dbapi_type = {} for bindparam in self.compiled.bind_names: key = self.compiled.bind_names[bindparam] - dialect_impl = bindparam.type.dialect_impl(self.dialect) + dialect_impl = bindparam.type._unwrapped_dialect_impl(self.dialect) dialect_impl_cls = type(dialect_impl) dbtype = dialect_impl.get_dbapi_type(self.dialect.dbapi) if dbtype is not None and ( diff --git a/lib/sqlalchemy/sql/type_api.py b/lib/sqlalchemy/sql/type_api.py index 1d1d6e089..6a323683b 100644 --- a/lib/sqlalchemy/sql/type_api.py +++ b/lib/sqlalchemy/sql/type_api.py @@ -445,6 +445,20 @@ class TypeEngine(Visitable): except KeyError: return self._dialect_info(dialect)['impl'] + def _unwrapped_dialect_impl(self, dialect): + """Return the 'unwrapped' dialect impl for this type. + + For a type that applies wrapping logic (e.g. TypeDecorator), give + us the real, actual dialect-level type that is used. + + This is used by TypeDecorator itself as well at least one case where + dialects need to check that a particular specific dialect-level + type is in use, within the :meth:`.DefaultDialect.set_input_sizes` + method. + + """ + return self.dialect_impl(dialect) + def _cached_literal_processor(self, dialect): """Return a dialect-specific literal processor for this type.""" try: @@ -922,7 +936,7 @@ class TypeDecorator(SchemaEventTarget, TypeEngine): # otherwise adapt the impl type, link # to a copy of this TypeDecorator and return # that. - typedesc = self.load_dialect_impl(dialect).dialect_impl(dialect) + typedesc = self._unwrapped_dialect_impl(dialect) tt = self.copy() if not isinstance(tt, self.__class__): raise AssertionError('Type object %s does not properly ' @@ -989,6 +1003,20 @@ class TypeDecorator(SchemaEventTarget, TypeEngine): """ return self.impl + def _unwrapped_dialect_impl(self, dialect): + """Return the 'unwrapped' dialect impl for this type. + + For a type that applies wrapping logic (e.g. TypeDecorator), give + us the real, actual dialect-level type that is used. + + This is used by TypeDecorator itself as well at least one case where + dialects need to check that a particular specific dialect-level + type is in use, within the :meth:`.DefaultDialect.set_input_sizes` + method. + + """ + return self.load_dialect_impl(dialect).dialect_impl(dialect) + def __getattr__(self, key): """Proxy all other undefined accessors to the underlying implementation.""" |
