diff --git a/.gitignore b/.gitignore index f7c4fa4..ebb9874 100644 --- a/.gitignore +++ b/.gitignore @@ -91,3 +91,8 @@ ENV/ # PyCharm project settings .idea .idea/ + +.venv/ +.pytest_cache/ +.mypy_cache/ + diff --git a/ACKNOWLEDGMENTS b/ACKNOWLEDGMENTS index 0483991..08c43e2 100644 --- a/ACKNOWLEDGMENTS +++ b/ACKNOWLEDGMENTS @@ -1,8 +1,10 @@ -THIS PROJECT IS DERIVED FROM THE FOLLOWING PROJECTS/FORKS: - - https://github.com/LocusEnergy/sqlalchemy-vertica-python +THIS PROJECT IS THE MAINLY DERIVED FROM startappdev repo: + - https://github.com/startappdev/sqlalchemy-vertica + +THIS PROJECT WAS ALSO DERIVED FROM THE FOLLOWING PROJECTS: + - https://github.com/zzzeek/sqlalchemy - https://github.com/bluelabsio/vertica-sqlalchemy - - https://github.com/dennisobrien/sqlalchemy-vertica-python - https://github.com/Eighty20/sqlalchemy-vertica-python - -THANKS TO ALL THE GREAT PEOPLE WHO'VE CONTRIBUTED TO THESE PROJECTS. + - https://github.com/LocusEnergy/sqlalchemy-vertica-python + - https://github.com/dennisobrien/sqlalchemy-vertica-python diff --git a/README.rst b/README.rst index 796720a..32c1dfc 100644 --- a/README.rst +++ b/README.rst @@ -1,35 +1,208 @@ sqlalchemy-vertica ================== -Vertica dialect for sqlalchemy. +Modern **Vertica Analytic Database** dialect for **SQLAlchemy 2.0+** with full support for **Async operations**, **Alembic migrations**, and modern Python (3.9 - 3.14+). -Forked from the `Vertica dialect for sqlalchemy using vertica-python `_. +.. image:: https://img.shields.io/badge/SQLAlchemy-2.0+-blue.svg + :target: https://www.sqlalchemy.org/ +.. image:: https://img.shields.io/badge/python-3.9+-blue.svg + :target: https://www.python.org/ +.. image:: https://img.shields.io/badge/Vertica-11--24+-green.svg + :target: https://www.vertica.com/ +.. image:: https://img.shields.io/badge/license-MIT-green.svg + :target: https://opensource.org/licenses/MIT -.. code-block:: python - import sqlalchemy as sa - import urllib - # for pyodbc connection - sa.create_engine('vertica+pyodbc:///?odbc_connect=%s' % (urllib.quote('DSN=dsn'),)) +Features +-------- - # for turbodbc connection - sa.create_engine('vertica+turbodbc:///?DSN=dsn') +* **Full SQLAlchemy 2.0+ Architecture**: Built on ``DefaultDialect`` with query caching (``supports_statement_cache = True``), 2.0 execution semantics, and parameter-bound reflection. +* **First-Class Async Engine Support**: Run queries asynchronously with ``create_async_engine()`` and ``AsyncSession`` via ``vertica+vertica_python_async://`` without blocking the asyncio event loop. +* **Alembic Migrations**: Native ``VerticaImpl`` integration with transactional DDL, type synonym resolution, and index no-op handling (since Vertica utilizes projections). +* **Multi-Driver Support**: + * ``vertica-python`` (Synchronous pure-Python DBAPI driver) + * ``vertica-python-async`` (Asynchronous DBAPI adapter for non-blocking asyncio / FastAPI apps) + * ``pyodbc`` (ODBC driver) + * ``turbodbc`` (High-speed ODBC driver for Arrow / NumPy / Pandas data workflows) +* **Rich Vertica Data Types**: + * Geospatial: ``GEOMETRY``, ``GEOGRAPHY`` + * Identifiers: native ``UUID`` + * Large objects: ``LONG VARCHAR``, ``LONG VARBINARY`` (up to 32MB) + * Complex types: ``ARRAY``, ``MAP``, ``ROW`` (Vertica 10+) + * Temporal: ``TIMESTAMPTZ``, ``TIMETZ``, ``INTERVAL`` +* **Complete Reflection**: Automatic introspection of schemas, tables, temp tables, views, view definitions, columns, primary keys, foreign keys, unique constraints, check constraints, table & column comments. - # for vertica-python connection - sa.create_engine('vertica+vertica_python://user:pwd@host:port/database') Installation ------------ -From PyPI: :: +Install from PyPI with your desired driver extras: + +.. code-block:: bash + + # Pure Python sync driver (recommended for sync applications) + pip install "sqlalchemy-vertica[vertica-python]" + + # Pure Python async driver (for AsyncEngine / FastAPI / asyncio) + pip install "sqlalchemy-vertica[asyncio]" + + # ODBC drivers + pip install "sqlalchemy-vertica[pyodbc]" + pip install "sqlalchemy-vertica[turbodbc]" + + # Alembic migrations support + pip install "sqlalchemy-vertica[alembic]" + + # Install all drivers and tools + pip install "sqlalchemy-vertica[all]" + + +Connection Strings +------------------ + +.. code-block:: python + + import sqlalchemy as sa + from sqlalchemy.ext.asyncio import create_async_engine + + # 1. Async (for FastAPI / asyncio applications) + async_engine = create_async_engine( + "vertica+vertica_python_async://user:pwd@host:5433/database?connection_timeout=10" + ) + + # 2. Sync vertica-python + engine = sa.create_engine( + "vertica+vertica_python://user:pwd@host:5433/database?connection_timeout=10" + ) + + # 3. PyODBC with connection string + engine_pyodbc = sa.create_engine( + "vertica+pyodbc:///?odbc_connect=DSN%3DVerticaDSN" + ) + + # 4. Turbodbc with DSN + engine_turbodbc = sa.create_engine( + "vertica+turbodbc:///?DSN=VerticaDSN" + ) + + +Quick Start +----------- + +Synchronous SQLAlchemy 2.0 +^^^^^^^^^^^^^^^^^^^^^^^^^^ + +.. code-block:: python + + from sqlalchemy import create_engine, text + + engine = create_engine("vertica+vertica_python://user:pwd@localhost:5433/mydb") + + with engine.connect() as conn: + result = conn.execute(text("SELECT version()")) + print(result.scalar()) + + # Transaction block + with engine.begin() as conn: + conn.execute( + text("INSERT INTO my_table (name) VALUES (:name)"), + {"name": "Alice"} + ) + + +Asynchronous SQLAlchemy 2.0 & FastAPI +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +.. code-block:: python + + import asyncio + from sqlalchemy import text + from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker + + async def main(): + engine = create_async_engine( + "vertica+vertica_python_async://user:pwd@localhost:5433/mydb", + pool_size=10, + ) + + async with engine.connect() as conn: + result = await conn.execute(text("SELECT 1")) + print(result.scalar()) + + # Using AsyncSession + session_factory = async_sessionmaker(engine, class_=AsyncSession) + async with session_factory() as session: + result = await session.execute(text("SELECT COUNT(*) FROM my_table")) + print("Count:", result.scalar()) + + await engine.dispose() + + asyncio.run(main()) + + +Alembic Migrations +------------------ + +In your Alembic ``env.py``, simply import ``sqlalchemy_vertica``: + +.. code-block:: python + + import sqlalchemy_vertica # Registers VerticaImpl plugin automatically + from alembic import context + + # configure context + context.configure( + connection=connection, + target_metadata=target_metadata, + transactional_ddl=True, + ) + +Vertica does not support traditional B-tree indexes (it utilizes projections). ``sqlalchemy-vertica`` treats index creation/dropping as safe no-ops in migrations to ensure multi-database migration scripts run seamlessly. + + +Custom Data Types +----------------- + +.. code-block:: python + + from sqlalchemy import Column, Integer, Table, MetaData + from sqlalchemy_vertica import ( + GEOMETRY, + GEOGRAPHY, + UUID, + LONG_VARCHAR, + ARRAY, + MAP, + ROW, + TIMESTAMPTZ, + ) + + metadata = MetaData() + + places = Table( + "places", + metadata, + Column("id", Integer, primary_key=True, autoincrement=True), + Column("guid", UUID, nullable=False), + Column("description", LONG_VARCHAR), + Column("location", GEOMETRY(srid=4326)), + Column("tags", ARRAY(LONG_VARCHAR)), + Column("metadata", MAP(LONG_VARCHAR, LONG_VARCHAR)), + Column("created_at", TIMESTAMPTZ), + ) + + +Testing & Coverage +------------------ + +Run the automated test suite with ``pytest`` and ``pytest-cov``: + +.. code-block:: bash - pip install sqlalchemy-vertica[pyodbc,turbodbc,vertica-python] # choose the relevant engines + pytest -v --cov=sqlalchemy_vertica --cov-report=term-missing -From git: :: - git clone https://github.com/startappdev/sqlalchemy-vertica - cd sqlalchemy-vertica - pip install pyodbc turbodbc vertica-python # choose the relevant engines - python setup.py install +License +------- +MIT License. See `LICENSE` for details. diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..50a5eb6 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,106 @@ +[build-system] +requires = ["setuptools>=61.0"] +build-backend = "setuptools.build_meta" + +[project] +name = "sqlalchemy-vertica" +version = "1.0.0" +description = "Vertica dialect for SQLAlchemy 2.0+ with Async & Alembic support" +readme = "README.rst" +license = "MIT" +authors = [ + {name = "Luis Villamarin", email = "luis@lv10.me"} +] +classifiers = [ + "Development Status :: 5 - Production/Stable", + "Intended Audience :: Developers", + "Programming Language :: Python :: 3", + "Programming Language :: Python :: 3.9", + "Programming Language :: Python :: 3.10", + "Programming Language :: Python :: 3.11", + "Programming Language :: Python :: 3.12", + "Programming Language :: Python :: 3.13", + "Programming Language :: Python :: 3.14", + "Topic :: Database", + "Topic :: Database :: Front-Ends", +] +requires-python = ">=3.9" +dependencies = [ + "sqlalchemy>=2.0.0", + "typing_extensions>=4.6.0", +] + +[project.optional-dependencies] +vertica-python = [ + "vertica-python>=1.0.0", +] +asyncio = [ + "vertica-python>=1.0.0", + "greenlet>=3.0.0", +] +pyodbc = [ + "pyodbc>=4.0.35", +] +turbodbc = [ + "turbodbc>=4.0.0", +] +alembic = [ + "alembic>=1.11.0", +] +all = [ + "vertica-python>=1.0.0", + "greenlet>=3.0.0", + "pyodbc>=4.0.35", + "turbodbc>=4.0.0", + "alembic>=1.11.0", +] +dev = [ + "pytest>=7.4.0", + "pytest-asyncio>=0.21.0", + "pytest-cov>=4.1.0", + "coverage[toml]>=7.3.0", + "mypy>=1.5.0", + "flake8>=6.0.0", +] + +[project.urls] +Homepage = "https://github.com/lv10/sqlalchemy-vertica" +Repository = "https://github.com/lv10/sqlalchemy-vertica" + +[project.entry-points."sqlalchemy.dialects"] +vertica = "sqlalchemy_vertica.dialect_vertica_python:VerticaDialect" +"vertica.vertica_python" = "sqlalchemy_vertica.dialect_vertica_python:VerticaDialect" +"vertica.vertica_python_async" = "sqlalchemy_vertica.dialect_vertica_python_async:VerticaDialect_vertica_python_async" +"vertica.async_vertica_python" = "sqlalchemy_vertica.dialect_vertica_python_async:VerticaDialect_vertica_python_async" +"vertica.pyodbc" = "sqlalchemy_vertica.dialect_pyodbc:VerticaDialect" +"vertica.turbodbc" = "sqlalchemy_vertica.dialect_turbodbc:VerticaDialect" + +[tool.setuptools.packages.find] +where = ["."] +include = ["sqlalchemy_vertica*"] + +[tool.pytest.ini_options] +minversion = "7.0" +addopts = "-ra --cov=sqlalchemy_vertica --cov-report=term-missing" +testpaths = ["tests"] +asyncio_mode = "auto" + +[tool.coverage.run] +source = ["sqlalchemy_vertica"] +branch = true + +[tool.coverage.report] +show_missing = true +skip_covered = false +exclude_lines = [ + "pragma: no cover", + "def __repr__", + "if TYPE_CHECKING:", + "raise NotImplementedError", + "\\.\\.\\.", +] + +[tool.mypy] +python_version = "3.12" +warn_unused_configs = true +ignore_missing_imports = true diff --git a/requirements.txt b/requirements.txt index 8941043..f216213 100644 --- a/requirements.txt +++ b/requirements.txt @@ -1,2 +1,2 @@ -six >= 1.10.0 -sqlalchemy >= 1.1.5 +sqlalchemy>=2.0.0 +typing_extensions>=4.6.0 diff --git a/setup.cfg b/setup.cfg deleted file mode 100644 index 5aef279..0000000 --- a/setup.cfg +++ /dev/null @@ -1,2 +0,0 @@ -[metadata] -description-file = README.rst diff --git a/setup.py b/setup.py index 9f34f33..6068493 100644 --- a/setup.py +++ b/setup.py @@ -1,43 +1,3 @@ from setuptools import setup -with open("README.rst", "r") as f: - description = f.read() - -version_info = (0, 2, 5) -version = '.'.join(map(str, version_info)) - -setup( - name='sqlalchemy-vertica', - version=version, - description='Vertica dialect for sqlalchemy', - long_description=description, - license='MIT', - url='https://github.com/startappdev/sqlalchemy-vertica', - download_url='https://github.com/startappdev/sqlalchemy-vertica/tarball/%s' % (version,), - author='StartApp Inc.', - author_email='ben.feinstein@startapp.com', - packages=( - 'sqlalchemy_vertica', - ), - install_requires=( - 'six >= 1.10.0', - 'sqlalchemy >= 1.1.11', - ), - extras_require={ - 'pyodbc': [ - 'pyodbc>=4.0.16', - ], - 'vertica-python': [ - 'psycopg2>=2.7.1', - 'vertica-python>=0.7.3', - ], - }, - entry_points={ - 'sqlalchemy.dialects': [ - 'vertica.pyodbc = ' - 'sqlalchemy_vertica.dialect_pyodbc:VerticaDialect [pyodbc]', - 'vertica.vertica_python = ' - 'sqlalchemy_vertica.dialect_vertica_python:VerticaDialect [vertica-python]', - ] - } -) +setup() diff --git a/sqlalchemy_vertica/__init__.py b/sqlalchemy_vertica/__init__.py index e69de29..3299d87 100644 --- a/sqlalchemy_vertica/__init__.py +++ b/sqlalchemy_vertica/__init__.py @@ -0,0 +1,78 @@ +from __future__ import annotations + +from sqlalchemy.dialects import registry + +from .base import ( + VerticaCompiler, + VerticaDDLCompiler, + VerticaDialect, + VerticaExecutionContext, + VerticaIdentifierPreparer, + VerticaTypeCompiler, +) +from .types import ( + ARRAY, + BYTEA, + DOUBLE_PRECISION, + GEOGRAPHY, + GEOMETRY, + INTERVAL, + LONG_VARBINARY, + LONG_VARCHAR, + MAP, + RAW, + ROW, + TIMESTAMPTZ, + TIMETZ, + UUID, + VARBINARY, +) + +__version__ = "1.0.0" + +# Register dialects with SQLAlchemy registry +registry.register("vertica", "sqlalchemy_vertica.dialect_vertica_python", "VerticaDialect") +registry.register("vertica.vertica_python", "sqlalchemy_vertica.dialect_vertica_python", "VerticaDialect") +registry.register( + "vertica.vertica_python_async", + "sqlalchemy_vertica.dialect_vertica_python_async", + "VerticaDialect_vertica_python_async", +) +registry.register( + "vertica.async_vertica_python", + "sqlalchemy_vertica.dialect_vertica_python_async", + "VerticaDialect_vertica_python_async", +) +registry.register("vertica.pyodbc", "sqlalchemy_vertica.dialect_pyodbc", "VerticaDialect") +registry.register("vertica.turbodbc", "sqlalchemy_vertica.dialect_turbodbc", "VerticaDialect") + +# Try loading alembic plugin if alembic is present +try: + from . import alembic # noqa: F401 +except Exception: # pragma: no cover + pass + +__all__ = [ + "__version__", + "VerticaDialect", + "VerticaCompiler", + "VerticaDDLCompiler", + "VerticaTypeCompiler", + "VerticaIdentifierPreparer", + "VerticaExecutionContext", + "ARRAY", + "MAP", + "ROW", + "UUID", + "GEOMETRY", + "GEOGRAPHY", + "LONG_VARCHAR", + "LONG_VARBINARY", + "TIMESTAMPTZ", + "TIMETZ", + "INTERVAL", + "BYTEA", + "RAW", + "VARBINARY", + "DOUBLE_PRECISION", +] diff --git a/sqlalchemy_vertica/alembic.py b/sqlalchemy_vertica/alembic.py new file mode 100644 index 0000000..25a6073 --- /dev/null +++ b/sqlalchemy_vertica/alembic.py @@ -0,0 +1,44 @@ +from __future__ import annotations + +import logging +from typing import Any + +log = logging.getLogger(__name__) + +try: + from alembic.ddl.impl import DefaultImpl + + class VerticaImpl(DefaultImpl): + __dialect__ = "vertica" + transactional_ddl = True + + type_synonyms = DefaultImpl.type_synonyms + ( + {"INT", "INTEGER", "INT8", "BIGINT"}, + {"FLOAT", "FLOAT8", "DOUBLE PRECISION", "REAL"}, + {"VARCHAR", "VARCHAR2", "TEXT", "LONG VARCHAR"}, + {"TIMESTAMP", "TIMESTAMP WITHOUT TIME ZONE", "DATETIME", "SMALLDATETIME"}, + {"TIMESTAMPTZ", "TIMESTAMP WITH TIME ZONE", "TIMESTAMP WITH TIMEZONE"}, + {"TIME", "TIME WITHOUT TIME ZONE"}, + {"TIMETZ", "TIME WITH TIME ZONE", "TIME WITH TIMEZONE"}, + {"VARBINARY", "BINARY", "RAW", "BYTEA", "LONG VARBINARY", "BLOB"}, + {"NUMERIC", "DECIMAL", "NUMBER", "MONEY"}, + {"BOOLEAN", "BOOL"}, + ) + + def create_index(self, index: Any, **kw: Any) -> None: + # Vertica is a columnar database using projections; it does not support indexes. + # We treat this as a safe no-op to allow generic migrations to succeed. + log.warning( + "Vertica does not support indexes (projections are used instead). " + "Ignoring create_index for '%s'.", + getattr(index, "name", "unnamed"), + ) + + def drop_index(self, index: Any, **kw: Any) -> None: + log.warning( + "Vertica does not support indexes. Ignoring drop_index for '%s'.", + getattr(index, "name", "unnamed"), + ) + +except ImportError: # pragma: no cover + VerticaImpl = None # type: ignore[misc,assignment] diff --git a/sqlalchemy_vertica/base.py b/sqlalchemy_vertica/base.py index b039a86..8477ad7 100644 --- a/sqlalchemy_vertica/base.py +++ b/sqlalchemy_vertica/base.py @@ -1,343 +1,863 @@ -from __future__ import absolute_import, unicode_literals, print_function, division +from __future__ import annotations +import itertools import re -from sqlalchemy import exc -from sqlalchemy import sql -from textwrap import dedent - -from sqlalchemy.dialects.postgresql import BYTEA, DOUBLE_PRECISION -from sqlalchemy.dialects.postgresql.base import PGDialect, PGDDLCompiler -from sqlalchemy.engine import reflection -from sqlalchemy.types import INTEGER, BIGINT, SMALLINT, VARCHAR, CHAR, \ - NUMERIC, FLOAT, REAL, DATE, DATETIME, BOOLEAN, BLOB, TIMESTAMP, TIME - -ischema_names = { - 'INT': INTEGER, - 'INTEGER': INTEGER, - 'INT8': INTEGER, - 'BIGINT': BIGINT, - 'SMALLINT': SMALLINT, - 'TINYINT': SMALLINT, - 'CHAR': CHAR, - 'VARCHAR': VARCHAR, - 'VARCHAR2': VARCHAR, - 'TEXT': VARCHAR, - 'NUMERIC': NUMERIC, - 'DECIMAL': NUMERIC, - 'NUMBER': NUMERIC, - 'MONEY': NUMERIC, - 'FLOAT': FLOAT, - 'FLOAT8': FLOAT, - 'REAL': REAL, - 'DOUBLE': DOUBLE_PRECISION, - 'TIMESTAMP': TIMESTAMP, - 'TIMESTAMP WITH TIMEZONE': TIMESTAMP, - 'TIME': TIME, - 'TIME WITH TIMEZONE': TIME, - 'DATE': DATE, - 'DATETIME': DATETIME, - 'SMALLDATETIME': DATETIME, - 'BINARY': BLOB, - 'VARBINARY': BLOB, - 'RAW': BLOB, - 'BYTEA': BYTEA, - 'BOOLEAN': BOOLEAN, +from typing import Any, Dict, List, Optional, Sequence, Tuple, cast + +from sqlalchemy import exc, sql, util +from sqlalchemy.engine import default, reflection +from sqlalchemy.engine.interfaces import ( + ReflectedCheckConstraint, + ReflectedColumn, + ReflectedForeignKeyConstraint, + ReflectedIndex, + ReflectedPrimaryKeyConstraint, + ReflectedTableComment, + ReflectedUniqueConstraint, +) +from sqlalchemy.sql import compiler, sqltypes +from sqlalchemy.types import ( + BIGINT, + BLOB, + BOOLEAN, + CHAR, + DATE, + DATETIME, + DECIMAL, + FLOAT, + INTEGER, + NUMERIC, + REAL, + SMALLINT, + VARCHAR, +) + +from .types import ( + ARRAY, + BYTEA, + DOUBLE_PRECISION, + GEOGRAPHY, + GEOMETRY, + INTERVAL, + LONG_VARBINARY, + LONG_VARCHAR, + MAP, + RAW, + ROW, + TIME, + TIMESTAMPTZ, + TIMETZ, + TIMESTAMP, + UUID, + VARBINARY, +) + +RESERVED_WORDS = { + "all", "analyze", "and", "any", "array", "as", "asc", "authorization", + "between", "bigint", "binary", "bit", "boolean", "both", "by", "case", + "cast", "char", "character", "check", "coalesce", "collate", "column", + "constraint", "correlation", "create", "cross", "current_date", + "current_time", "current_timestamp", "current_user", "date", "decimal", + "default", "deferrable", "deferred", "delete", "desc", "direct", + "distinct", "do", "double", "drop", "else", "end", "except", "exists", + "false", "float", "float8", "for", "foreign", "freeze", "from", "full", + "function", "grant", "group", "having", "ilike", "in", "initially", + "inner", "inout", "insert", "instead", "int", "integer", "intersect", + "interval", "into", "is", "isnull", "join", "ksafe", "leading", "left", + "like", "limit", "localtime", "localtimestamp", "long", "match", + "natural", "new", "not", "notnull", "null", "nullif", "number", + "numeric", "off", "offset", "old", "on", "only", "or", "order", "out", + "outer", "over", "overlaps", "partition", "placing", "precision", + "primary", "projection", "raw", "real", "references", "rejected", + "rename", "replace", "right", "row", "schema", "select", "session_user", + "set", "setof", "similar", "smallint", "some", "substring", "table", + "then", "time", "timestamp", "timestamptz", "timetz", "tinyint", "to", + "trailing", "treat", "true", "truncate", "uncommitted", "union", "unique", + "unknown", "unsegmented", "update", "user", "using", "uuid", "values", + "varbinary", "varchar", "varchar2", "varying", "view", "when", "where", + "with", } +ischema_names: Dict[str, Any] = { + "INT": INTEGER, + "INTEGER": INTEGER, + "INT8": INTEGER, + "BIGINT": BIGINT, + "SMALLINT": SMALLINT, + "TINYINT": SMALLINT, + "CHAR": CHAR, + "VARCHAR": VARCHAR, + "VARCHAR2": VARCHAR, + "TEXT": VARCHAR, + "LONG VARCHAR": LONG_VARCHAR, + "NUMERIC": NUMERIC, + "DECIMAL": DECIMAL, + "NUMBER": NUMERIC, + "MONEY": NUMERIC, + "FLOAT": FLOAT, + "FLOAT8": FLOAT, + "REAL": REAL, + "DOUBLE": DOUBLE_PRECISION, + "DOUBLE PRECISION": DOUBLE_PRECISION, + "TIMESTAMP": TIMESTAMP, + "TIMESTAMP WITHOUT TIME ZONE": TIMESTAMP, + "TIMESTAMP WITH TIMEZONE": TIMESTAMPTZ, + "TIMESTAMP WITH TIME ZONE": TIMESTAMPTZ, + "TIMESTAMPTZ": TIMESTAMPTZ, + "TIME": TIME, + "TIME WITHOUT TIME ZONE": TIME, + "TIME WITH TIMEZONE": TIMETZ, + "TIME WITH TIME ZONE": TIMETZ, + "TIMETZ": TIMETZ, + "INTERVAL": INTERVAL, + "INTERVAL DAY": INTERVAL, + "INTERVAL DAY TO SECOND": INTERVAL, + "INTERVAL YEAR TO MONTH": INTERVAL, + "DATE": DATE, + "DATETIME": DATETIME, + "SMALLDATETIME": DATETIME, + "BINARY": VARBINARY, + "VARBINARY": VARBINARY, + "RAW": RAW, + "BYTEA": BYTEA, + "BLOB": BLOB, + "BOOLEAN": BOOLEAN, + "BOOL": BOOLEAN, + "LONG VARBINARY": LONG_VARBINARY, + "GEOMETRY": GEOMETRY, + "GEOGRAPHY": GEOGRAPHY, + "UUID": UUID, + "ARRAY": ARRAY, + "MAP": MAP, + "ROW": ROW, +} + + +class VerticaIdentifierPreparer(compiler.IdentifierPreparer): + reserved_words = RESERVED_WORDS + + def __init__(self, dialect: default.DefaultDialect, **kw: Any) -> None: + super().__init__( + dialect, + initial_quote='"', + final_quote='"', + **kw, + ) + -class VerticaDDLCompiler(PGDDLCompiler): - def get_column_specification(self, column, **kwargs): +class VerticaTypeCompiler(compiler.GenericTypeCompiler): + def visit_INTEGER(self, type_: Any, **kw: Any) -> str: + return "INT" + + def visit_BIGINT(self, type_: Any, **kw: Any) -> str: + return "BIGINT" + + def visit_SMALLINT(self, type_: Any, **kw: Any) -> str: + return "SMALLINT" + + def visit_FLOAT(self, type_: Any, **kw: Any) -> str: + if type_.precision is not None: + return f"FLOAT({type_.precision})" + return "FLOAT" + + def visit_DOUBLE_PRECISION(self, type_: Any, **kw: Any) -> str: + return "DOUBLE PRECISION" + + def visit_REAL(self, type_: Any, **kw: Any) -> str: + return "REAL" + + def visit_NUMERIC(self, type_: Any, **kw: Any) -> str: + if type_.precision is None: + return "NUMERIC" + if type_.scale is None: + return f"NUMERIC({type_.precision})" + return f"NUMERIC({type_.precision}, {type_.scale})" + + def visit_DECIMAL(self, type_: Any, **kw: Any) -> str: + return self.visit_NUMERIC(type_, **kw) + + def visit_VARCHAR(self, type_: Any, **kw: Any) -> str: + if type_.length: + return f"VARCHAR({type_.length})" + return "VARCHAR" + + def visit_CHAR(self, type_: Any, **kw: Any) -> str: + if type_.length: + return f"CHAR({type_.length})" + return "CHAR" + + def visit_TEXT(self, type_: Any, **kw: Any) -> str: + return "LONG VARCHAR" + + def visit_LONG_VARCHAR(self, type_: Any, **kw: Any) -> str: + if type_.length: + return f"LONG VARCHAR({type_.length})" + return "LONG VARCHAR" + + def visit_BLOB(self, type_: Any, **kw: Any) -> str: + return "LONG VARBINARY" + + def visit_VARBINARY(self, type_: Any, **kw: Any) -> str: + if type_.length: + return f"VARBINARY({type_.length})" + return "VARBINARY" + + def visit_LONG_VARBINARY(self, type_: Any, **kw: Any) -> str: + if type_.length: + return f"LONG VARBINARY({type_.length})" + return "LONG VARBINARY" + + def visit_large_binary(self, type_: Any, **kw: Any) -> str: + if getattr(type_, "length", None): + return f"VARBINARY({type_.length})" + return "VARBINARY" + + def visit_BYTEA(self, type_: Any, **kw: Any) -> str: + return "BYTEA" + + def visit_RAW(self, type_: Any, **kw: Any) -> str: + return "RAW" + + def visit_BOOLEAN(self, type_: Any, **kw: Any) -> str: + return "BOOLEAN" + + def visit_DATE(self, type_: Any, **kw: Any) -> str: + return "DATE" + + def visit_TIME(self, type_: Any, **kw: Any) -> str: + if getattr(type_, "timezone", False): + return self.visit_TIMETZ(type_, **kw) + if getattr(type_, "precision", None) is not None: + return f"TIME({type_.precision})" + return "TIME" + + def visit_TIMETZ(self, type_: Any, **kw: Any) -> str: + if getattr(type_, "precision", None) is not None: + return f"TIMETZ({type_.precision})" + return "TIMETZ" + + def visit_TIMESTAMP(self, type_: Any, **kw: Any) -> str: + if getattr(type_, "timezone", False): + return self.visit_TIMESTAMPTZ(type_, **kw) + if getattr(type_, "precision", None) is not None: + return f"TIMESTAMP({type_.precision})" + return "TIMESTAMP" + + def visit_TIMESTAMPTZ(self, type_: Any, **kw: Any) -> str: + if getattr(type_, "precision", None) is not None: + return f"TIMESTAMPTZ({type_.precision})" + return "TIMESTAMPTZ" + + def visit_DATETIME(self, type_: Any, **kw: Any) -> str: + return "DATETIME" + + def visit_INTERVAL(self, type_: Any, **kw: Any) -> str: + parts = ["INTERVAL"] + if getattr(type_, "fields", None): + parts.append(str(type_.fields)) + if getattr(type_, "precision", None) is not None: + parts.append(f"({type_.precision})") + return " ".join(parts) + + def visit_UUID(self, type_: Any, **kw: Any) -> str: + return "UUID" + + def visit_GEOMETRY(self, type_: Any, **kw: Any) -> str: + if getattr(type_, "srid", None) is not None: + return f"GEOMETRY({type_.srid})" + return "GEOMETRY" + + def visit_GEOGRAPHY(self, type_: Any, **kw: Any) -> str: + if getattr(type_, "srid", None) is not None: + return f"GEOGRAPHY({type_.srid})" + return "GEOGRAPHY" + + def visit_ARRAY(self, type_: Any, **kw: Any) -> str: + inner = self.process(type_.item_type, **kw) + if getattr(type_, "length", None) is not None: + return f"ARRAY[{inner}, {type_.length}]" + return f"ARRAY[{inner}]" + + def visit_MAP(self, type_: Any, **kw: Any) -> str: + k = self.process(type_.key_type, **kw) + v = self.process(type_.value_type, **kw) + return f"MAP[{k}, {v}]" + + def visit_ROW(self, type_: Any, **kw: Any) -> str: + fields = [ + f"{self.dialect.identifier_preparer.quote(name)} {self.process(ftype, **kw)}" + for name, ftype in type_.fields.items() + ] + return f"ROW({', '.join(fields)})" + + +class VerticaCompiler(compiler.SQLCompiler): + def visit_sequence(self, sequence: Any, **kw: Any) -> str: + seq_name = self.preparer.format_sequence(sequence) + return f"{seq_name}.NEXTVAL" + + def limit_clause(self, select: Any, **kw: Any) -> str: + text = "" + if select._limit_clause is not None: + text += f" \n LIMIT {self.process(select._limit_clause, **kw)}" + if select._offset_clause is not None: + text += f" OFFSET {self.process(select._offset_clause, **kw)}" + return text + + def for_update_clause(self, select: Any, **kw: Any) -> str: + return " FOR UPDATE" + + +class VerticaDDLCompiler(compiler.DDLCompiler): + def get_column_specification(self, column: Any, **kwargs: Any) -> str: colspec = self.preparer.format_column(column) - # noinspection PyUnusedLocal - impl_type = column.type.dialect_impl(self.dialect) - # noinspection PyProtectedMember + colspec += " " + self.dialect.type_compiler.process(column.type) + if column.primary_key and column is column.table._autoincrement_column: colspec += " AUTO_INCREMENT" else: - colspec += " " + self.dialect.type_compiler.process(column.type) - default = self.get_column_default_string(column) - if default is not None: - colspec += " DEFAULT " + default + default_str = self.get_column_default_string(column) + if default_str is not None: + colspec += " DEFAULT " + default_str if not column.nullable: colspec += " NOT NULL" - return colspec + return colspec -# noinspection PyArgumentList,PyAbstractClass -class VerticaDialect(PGDialect): - name = 'vertica' - ischema_names = ischema_names + def visit_create_index( + self, + create: Any, + include_schema: bool = False, + include_table_schema: bool = False, + **kw: Any, + ) -> str: + # Vertica is a columnar database using projections; it does not support indexes. + return "" + + def visit_drop_index(self, drop: Any, **kw: Any) -> str: + return "" + + def visit_primary_key_constraint(self, constraint: Any, **kw: Any) -> str: + cols = ", ".join(self.preparer.quote(c.name) for c in constraint.columns) + text = f"PRIMARY KEY ({cols})" + if constraint.name: + text = f"CONSTRAINT {self.preparer.quote(constraint.name)} {text}" + return text + + def visit_foreign_key_constraint(self, constraint: Any, **kw: Any) -> str: + cols = ", ".join(self.preparer.quote(c.name) for c in constraint.columns) + ref_cols = ", ".join(self.preparer.quote(elem.column.name) for elem in constraint.elements) + ref_table = self.preparer.format_table(constraint.referred_table) + text = f"FOREIGN KEY ({cols}) REFERENCES {ref_table} ({ref_cols})" + if constraint.name: + text = f"CONSTRAINT {self.preparer.quote(constraint.name)} {text}" + return text + + def visit_unique_constraint(self, constraint: Any, **kw: Any) -> str: + cols = ", ".join(self.preparer.quote(c.name) for c in constraint.columns) + text = f"UNIQUE ({cols})" + if constraint.name: + text = f"CONSTRAINT {self.preparer.quote(constraint.name)} {text}" + return text + + def visit_check_constraint(self, constraint: Any, **kw: Any) -> str: + text = f"CHECK ({self.sql_compiler.process(constraint.sqltext, include_table=False)})" + if constraint.name: + text = f"CONSTRAINT {self.preparer.quote(constraint.name)} {text}" + return text + + +class VerticaExecutionContext(default.DefaultExecutionContext): + pass + + +class VerticaDialect(default.DefaultDialect): + name = "vertica" + supports_statement_cache = True + + # Feature capabilities in Vertica + supports_native_boolean = True + supports_native_decimal = True + supports_native_uuid = True + supports_alter = True + supports_sequences = True + supports_identity_columns = True + supports_comments = True + supports_schemas = True + supports_views = True + supports_multivalues_insert = True + insert_returning = False + update_returning = False + delete_returning = False + use_insertmanyvalues = False + postfetch_lastrowid = False + default_schema_name = "public" + + # Dialect components + statement_compiler = VerticaCompiler ddl_compiler = VerticaDDLCompiler + type_compiler_cls = VerticaTypeCompiler + preparer: type[compiler.IdentifierPreparer] = VerticaIdentifierPreparer + execution_ctx_cls = VerticaExecutionContext + ischema_names = ischema_names - def _get_default_schema_name(self, connection): - return connection.scalar("SELECT current_schema()") - - def _get_server_version_info(self, connection): - v = connection.scalar("SELECT version()") - m = re.match(r".*Vertica Analytic Database v(\d+)\.(\d+)\.(\d)+.*", v) - if not m: - raise AssertionError("Could not determine version from string '%(ver)s'" % {'ver': v}) - return tuple([int(x) for x in m.group(1, 2, 3) if x is not None]) - - # noinspection PyRedeclaration - def _get_default_schema_name(self, connection): - return connection.scalar("SELECT current_schema()") - - def create_connect_args(self, url): - opts = url.translate_connect_args(username='user') + def _get_default_schema_name(self, connection: Any) -> str: + schema = connection.scalar(sql.text("SELECT current_schema()")) + return str(schema) if schema else "public" + + def _get_server_version_info(self, connection: Any) -> Tuple[int, ...]: + v = connection.scalar(sql.text("SELECT version()")) + if not v: + return (0, 0, 0) + m = re.search(r"v(\d+)\.(\d+)(?:\.(\d+))?", str(v)) + if m: + groups = [int(x) for x in m.groups() if x is not None] + while len(groups) < 3: + groups.append(0) + return tuple(groups) + return (0, 0, 0) + + def create_connect_args(self, url: Any) -> Tuple[Sequence[Any], Dict[str, Any]]: + opts = url.translate_connect_args(username="user") opts.update(url.query) return [], opts - def has_schema(self, connection, schema): - has_schema_sql = sql.text(dedent(""" - SELECT EXISTS ( - SELECT schema_name - FROM v_catalog.schemata - WHERE lower(schema_name) = '%(schema)s') - """ % {'schema': schema.lower()})) - - c = connection.execute(has_schema_sql) - return bool(c.scalar()) - - def has_table(self, connection, table_name, schema=None): + def has_schema(self, connection: Any, schema_name: str, **kw: Any) -> bool: + stmt = sql.text( + "SELECT EXISTS (" + " SELECT schema_name FROM v_catalog.schemata " + " WHERE lower(schema_name) = lower(:schema)" + ")" + ) + res = connection.scalar(stmt, {"schema": schema_name}) + return bool(res) + + def has_table( + self, connection: Any, table_name: str, schema: Optional[str] = None, **kw: Any + ) -> bool: if schema is None: schema = self._get_default_schema_name(connection) - has_table_sql = sql.text(dedent(""" - SELECT EXISTS ( - SELECT table_name - FROM v_catalog.all_tables - WHERE lower(table_name) = '%(table)s' - AND lower(schema_name) = '%(schema)s') - """ % {'schema': schema.lower(), 'table': table_name.lower()})) - - c = connection.execute(has_table_sql) - return bool(c.scalar()) - - def has_sequence(self, connection, sequence_name, schema=None): + stmt = sql.text( + "SELECT EXISTS (" + " SELECT table_name FROM v_catalog.all_tables " + " WHERE lower(table_name) = lower(:table) " + " AND lower(schema_name) = lower(:schema)" + ")" + ) + res = connection.scalar(stmt, {"table": table_name, "schema": schema}) + return bool(res) + + def has_sequence( + self, connection: Any, sequence_name: str, schema: Optional[str] = None, **kw: Any + ) -> bool: if schema is None: schema = self._get_default_schema_name(connection) - has_seq_sql = sql.text(dedent(""" - SELECT EXISTS ( - SELECT sequence_name - FROM v_catalog.sequences - WHERE lower(sequence_name) = '%(sequence)s' - AND lower(sequence_schema) = '%(schema)s') - """ % {'schema': schema.lower(), 'sequence': sequence_name.lower()})) - - c = connection.execute(has_seq_sql) - return bool(c.scalar()) - - def has_type(self, connection, type_name, schema=None): - has_type_sql = sql.text(dedent(""" - SELECT EXISTS ( - SELECT type_name - FROM v_catalog.types - WHERE lower(type_name) = '%(type)s') - """ % {'type': type_name.lower()})) - - c = connection.execute(has_type_sql) - return bool(c.scalar()) - - @reflection.cache - def get_schema_names(self, connection, **kw): - get_schemas_sql = sql.text(dedent(""" - SELECT schema_name - FROM v_catalog.schemata - """)) - - c = connection.execute(get_schemas_sql) - return [row[0] for row in c if not row[0].startswith('v_')] + stmt = sql.text( + "SELECT EXISTS (" + " SELECT sequence_name FROM v_catalog.sequences " + " WHERE lower(sequence_name) = lower(:sequence) " + " AND lower(sequence_schema) = lower(:schema)" + ")" + ) + res = connection.scalar(stmt, {"sequence": sequence_name, "schema": schema}) + return bool(res) + + def has_type( + self, connection: Any, type_name: str, schema: Optional[str] = None, **kw: Any + ) -> bool: + stmt = sql.text( + "SELECT EXISTS (" + " SELECT type_name FROM v_catalog.types " + " WHERE lower(type_name) = lower(:type)" + ")" + ) + res = connection.scalar(stmt, {"type": type_name}) + return bool(res) @reflection.cache - def get_table_oid(self, connection, table_name, schema=None, **kw): - if schema is None: - schema = self._get_default_schema_name(connection) - - get_oid_sql = sql.text(dedent(""" - SELECT table_id - FROM v_catalog.tables - WHERE lower(table_name) = '%(table)s' - AND lower(table_schema) = '%(schema)s' - """ % {'schema': schema.lower(), 'table': table_name.lower()})) - - c = connection.execute(get_oid_sql) - table_oid = c.scalar() - if table_oid is None: - raise exc.NoSuchTableError(table_name) - return table_oid + def get_schema_names(self, connection: Any, **kw: Any) -> List[str]: + stmt = sql.text( + "SELECT schema_name FROM v_catalog.schemata " + "ORDER BY schema_name" + ) + rows = connection.execute(stmt).fetchall() + system_schemas = {"v_catalog", "v_monitor", "v_internal", "v_txtindex", "txtindex"} + return [row[0] for row in rows if row[0].lower() not in system_schemas] @reflection.cache - def get_table_names(self, connection, schema=None, **kw): + def get_table_names( + self, connection: Any, schema: Optional[str] = None, **kw: Any + ) -> List[str]: if schema is not None: - schema_condition = "lower(table_schema) = '%(schema)s'" % {'schema': schema.lower()} + stmt = sql.text( + "SELECT table_name FROM v_catalog.tables " + "WHERE lower(table_schema) = lower(:schema) " + " AND NOT is_system_table " + "ORDER BY table_name" + ) + rows = connection.execute(stmt, {"schema": schema}).fetchall() else: - schema_condition = "1" - - get_tables_sql = sql.text(dedent(""" - SELECT table_name - FROM v_catalog.tables - WHERE %(schema_condition)s - ORDER BY table_schema, table_name - """ % {'schema_condition': schema_condition})) - - c = connection.execute(get_tables_sql) - return [row[0] for row in c] + stmt = sql.text( + "SELECT table_name FROM v_catalog.tables " + "WHERE NOT is_system_table " + "ORDER BY table_schema, table_name" + ) + rows = connection.execute(stmt).fetchall() + return [row[0] for row in rows] @reflection.cache - def get_temp_table_names(self, connection, schema=None, **kw): + def get_temp_table_names( + self, connection: Any, schema: Optional[str] = None, **kw: Any + ) -> List[str]: if schema is not None: - schema_condition = "lower(table_schema) = '%(schema)s'" % {'schema': schema.lower()} + stmt = sql.text( + "SELECT table_name FROM v_catalog.tables " + "WHERE lower(table_schema) = lower(:schema) " + " AND is_temp_table " + "ORDER BY table_name" + ) + rows = connection.execute(stmt, {"schema": schema}).fetchall() else: - schema_condition = "1" - - get_tables_sql = sql.text(dedent(""" - SELECT table_name - FROM v_catalog.tables - WHERE %(schema_condition)s - AND IS_TEMP_TABLE - ORDER BY table_schema, table_name - """ % {'schema_condition': schema_condition})) - - c = connection.execute(get_tables_sql) - return [row[0] for row in c] + stmt = sql.text( + "SELECT table_name FROM v_catalog.tables " + "WHERE is_temp_table " + "ORDER BY table_schema, table_name" + ) + rows = connection.execute(stmt).fetchall() + return [row[0] for row in rows] @reflection.cache - def get_view_names(self, connection, schema=None, **kw): + def get_view_names( + self, connection: Any, schema: Optional[str] = None, **kw: Any + ) -> List[str]: if schema is not None: - schema_condition = "lower(table_schema) = '%(schema)s'" % {'schema': schema.lower()} + stmt = sql.text( + "SELECT table_name FROM v_catalog.views " + "WHERE lower(table_schema) = lower(:schema) " + " AND NOT is_system_view " + "ORDER BY table_name" + ) + rows = connection.execute(stmt, {"schema": schema}).fetchall() else: - schema_condition = "1" + stmt = sql.text( + "SELECT table_name FROM v_catalog.views " + "WHERE NOT is_system_view " + "ORDER BY table_schema, table_name" + ) + rows = connection.execute(stmt).fetchall() + return [row[0] for row in rows] - get_views_sql = sql.text(dedent(""" - SELECT table_name - FROM v_catalog.views - WHERE %(schema_condition)s - ORDER BY table_schema, table_name - """ % {'schema_condition': schema_condition})) + @reflection.cache + def get_view_definition( + self, + connection: Any, + view_name: str, + schema: Optional[str] = None, + **kw: Any, + ) -> str: + if schema is None: + schema = self._get_default_schema_name(connection) - c = connection.execute(get_views_sql) - return [row[0] for row in c] + stmt = sql.text( + "SELECT view_definition FROM v_catalog.views " + "WHERE lower(table_name) = lower(:table) " + " AND lower(table_schema) = lower(:schema)" + ) + res = connection.scalar(stmt, {"table": view_name, "schema": schema}) + return str(res) if res is not None else "" @reflection.cache - def get_temp_view_names(self, connection, schema=None, **kw): - return [] + def get_table_comment( + self, connection: Any, table_name: str, schema: Optional[str] = None, **kw: Any + ) -> ReflectedTableComment: + if schema is None: + schema = self._get_default_schema_name(connection) - @reflection.cache - def get_columns(self, connection, table_name, schema=None, **kw): - if schema is not None: - schema_condition = "lower(table_schema) = '%(schema)s'" % {'schema': schema.lower()} - else: - schema_condition = "1" - - s = sql.text(dedent(""" - SELECT column_name, data_type, column_default, is_nullable - FROM v_catalog.columns - WHERE lower(table_name) = '%(table)s' - AND %(schema_condition)s - UNION ALL - SELECT column_name, data_type, '' as column_default, true as is_nullable - FROM v_catalog.view_columns - WHERE lower(table_name) = '%(table)s' - AND %(schema_condition)s - """ % {'table': table_name.lower(), 'schema_condition': schema_condition})) - - spk = sql.text(dedent(""" - SELECT column_name - FROM v_catalog.primary_keys - WHERE lower(table_name) = '%(table)s' - AND constraint_type = 'p' - AND %(schema_condition)s - """ % {'table': table_name.lower(), 'schema_condition': schema_condition})) - - pk_columns = [x[0] for x in connection.execute(spk)] - columns = [] - for row in connection.execute(s): - name = row.column_name - dtype = row.data_type.upper() - if '(' in dtype: - dtype = dtype.split('(')[0] - coltype = self.ischema_names[dtype] - primary_key = name in pk_columns - default = row.column_default - nullable = row.is_nullable - - columns.append({ - 'name': name, - 'type': coltype, - 'nullable': nullable, - 'default': default, - 'primary_key': primary_key - }) - return columns + stmt = sql.text( + "SELECT comment FROM v_catalog.comments " + "WHERE object_type = 'TABLE' " + " AND lower(object_name) = lower(:table) " + " AND lower(object_schema) = lower(:schema)" + ) + res = connection.scalar(stmt, {"table": table_name, "schema": schema}) + return cast(ReflectedTableComment, {"text": str(res) if res is not None else None}) @reflection.cache - def get_unique_constraints(self, connection, table_name, schema=None, **kw): + def get_table_oid( + self, connection: Any, table_name: str, schema: Optional[str] = None, **kw: Any + ) -> int: if schema is None: schema = self._get_default_schema_name(connection) - get_constrains_sql = sql.text(dedent(""" - SELECT constraint_name, column_name - FROM v_catalog.constraint_columns - WHERE lower(table_name) = '%(table)s' - -- AND constraint_type IN ('p', 'u') - AND lower(table_schema) = '%(schema)s' - """ % {'schema': schema.lower(), 'table': table_name.lower()})) + stmt = sql.text( + "SELECT table_id FROM (" + " SELECT table_id, table_name, table_schema FROM v_catalog.tables " + " UNION " + " SELECT table_id, table_name, table_schema FROM v_catalog.views" + ") AS a " + "WHERE lower(a.table_name) = lower(:table) " + " AND lower(a.table_schema) = lower(:schema)" + ) + table_oid = connection.scalar(stmt, {"table": table_name, "schema": schema}) + if table_oid is None: + raise exc.NoSuchTableError(f"{schema}.{table_name}" if schema else table_name) + return int(table_oid) - c = connection.execute(get_constrains_sql) - if c.rowcount <= 0: - return [] + @reflection.cache + def get_columns( + self, connection: Any, table_name: str, schema: Optional[str] = None, **kw: Any + ) -> List[ReflectedColumn]: + if schema is None: + schema = self._get_default_schema_name(connection) - constraints, columns = zip(*c) - result_dict = { - unique_con: [col for con, col in zip(constraints, columns) if con == unique_con] - for unique_con in set(constraints) - } + cols_stmt = sql.text( + "SELECT column_name, data_type, column_default, is_nullable, ordinal_position " + "FROM v_catalog.columns " + "WHERE lower(table_name) = lower(:table) " + " AND lower(table_schema) = lower(:schema) " + "UNION ALL " + "SELECT column_name, data_type, '' as column_default, true as is_nullable, ordinal_position " + "FROM v_catalog.view_columns " + "WHERE lower(table_name) = lower(:table) " + " AND lower(table_schema) = lower(:schema) " + "ORDER BY ordinal_position" + ) + cols_rows = connection.execute(cols_stmt, {"table": table_name, "schema": schema}).fetchall() + + if not cols_rows: + if not self.has_table(connection, table_name, schema=schema): + raise exc.NoSuchTableError(f"{schema}.{table_name}" if schema else table_name) + + pk_stmt = sql.text( + "SELECT column_name FROM v_catalog.primary_keys " + "WHERE lower(table_name) = lower(:table) " + " AND lower(table_schema) = lower(:schema)" + ) + pk_columns = {row[0].lower() for row in connection.execute(pk_stmt, {"table": table_name, "schema": schema})} + + comment_stmt = sql.text( + "SELECT sub_object_name, comment FROM v_catalog.comments " + "WHERE object_type = 'COLUMN' " + " AND lower(object_name) = lower(:table) " + " AND lower(object_schema) = lower(:schema)" + ) + col_comments = { + row[0].lower(): row[1] + for row in connection.execute(comment_stmt, {"table": table_name, "schema": schema}) + if row[0] is not None + } + + columns: List[ReflectedColumn] = [] + for row in cols_rows: + name = str(row[0]) + dtype = str(row[1]).lower() + default_val = row[2] if row[2] != "" else None + is_nullable = bool(row[3]) + primary_key = name.lower() in pk_columns + comment = col_comments.get(name.lower()) + + col_info = self._get_column_info( + name=name, + format_type=dtype, + default=default_val, + nullable=is_nullable, + schema=schema, + ) + col_info["primary_key"] = primary_key + col_info["comment"] = comment + columns.append(cast(ReflectedColumn, col_info)) - return [{"name": name, "column_names": cols} for name, cols in result_dict.items()] + return columns @reflection.cache - def get_check_constraints( - self, connection, table_name, schema=None, **kw): - table_oid = self.get_table_oid(connection, table_name, schema, - info_cache=kw.get('info_cache')) + def get_pk_constraint( + self, connection: Any, table_name: str, schema: Optional[str] = None, **kw: Any + ) -> ReflectedPrimaryKeyConstraint: + if schema is None: + schema = self._get_default_schema_name(connection) - constraints_sql = sql.text(dedent(""" - SELECT constraint_name, column_name - FROM v_catalog.constraint_columns - WHERE table_id = %(oid)s - AND constraint_type = 'c' - """ % {'oid': table_oid})) + stmt = sql.text( + "SELECT constraint_name, column_name FROM v_catalog.primary_keys " + "WHERE lower(table_name) = lower(:table) " + " AND lower(table_schema) = lower(:schema) " + "ORDER BY ordinal_position" + ) + rows = connection.execute(stmt, {"table": table_name, "schema": schema}).fetchall() - c = connection.execute(constraints_sql) + if not rows: + return cast(ReflectedPrimaryKeyConstraint, {"constrained_columns": [], "name": None}) - return [{'name': name, 'sqltext': col} for name, col in c.fetchall()] + constraint_name = rows[0][0] + constrained_columns = [row[1] for row in rows] + return cast( + ReflectedPrimaryKeyConstraint, + {"constrained_columns": constrained_columns, "name": constraint_name}, + ) - def normalize_name(self, name): - name = name and name.rstrip() - if name is None: - return None - return name.lower() + @reflection.cache + def get_foreign_keys( + self, connection: Any, table_name: str, schema: Optional[str] = None, **kw: Any + ) -> List[ReflectedForeignKeyConstraint]: + if schema is None: + schema = self._get_default_schema_name(connection) - def denormalize_name(self, name): - return name + stmt = sql.text( + "SELECT constraint_name, column_name, reference_table_schema, " + " reference_table_name, reference_column_name " + "FROM v_catalog.foreign_keys " + "WHERE lower(table_name) = lower(:table) " + " AND lower(table_schema) = lower(:schema) " + "ORDER BY constraint_name, ordinal_position" + ) + rows = connection.execute(stmt, {"table": table_name, "schema": schema}).fetchall() + + fkeys: List[ReflectedForeignKeyConstraint] = [] + for name, group in itertools.groupby(rows, key=lambda r: r[0]): + items = list(group) + fkeys.append( + cast( + ReflectedForeignKeyConstraint, + { + "name": name, + "constrained_columns": [item[1] for item in items], + "referred_schema": items[0][2], + "referred_table": items[0][3], + "referred_columns": [item[4] for item in items], + }, + ) + ) + return fkeys - # methods allows table introspection to work @reflection.cache - def get_pk_constraint(self, bind, table_name, schema=None, **kw): - return {'constrained_columns': [], 'name': 'undefined'} + def get_unique_constraints( + self, connection: Any, table_name: str, schema: Optional[str] = None, **kw: Any + ) -> List[ReflectedUniqueConstraint]: + if schema is None: + schema = self._get_default_schema_name(connection) + + stmt = sql.text( + "SELECT constraint_name, column_name FROM v_catalog.constraint_columns " + "WHERE lower(table_name) = lower(:table) " + " AND lower(table_schema) = lower(:schema) " + " AND constraint_type = 'u' " + "ORDER BY constraint_name, ordinal_position" + ) + rows = connection.execute(stmt, {"table": table_name, "schema": schema}).fetchall() + + constraints: List[ReflectedUniqueConstraint] = [] + for name, group in itertools.groupby(rows, key=lambda r: r[0]): + constraints.append( + cast( + ReflectedUniqueConstraint, + { + "name": name, + "column_names": [item[1] for item in group], + }, + ) + ) + return constraints @reflection.cache - def get_foreign_keys(self, connection, table_name, schema=None, **kw): - return [] + def get_check_constraints( + self, connection: Any, table_name: str, schema: Optional[str] = None, **kw: Any + ) -> List[ReflectedCheckConstraint]: + if schema is None: + schema = self._get_default_schema_name(connection) + + stmt = sql.text( + "SELECT constraint_name, column_name FROM v_catalog.constraint_columns " + "WHERE lower(table_name) = lower(:table) " + " AND lower(table_schema) = lower(:schema) " + " AND constraint_type = 'c' " + "ORDER BY constraint_name" + ) + rows = connection.execute(stmt, {"table": table_name, "schema": schema}).fetchall() + + return [cast(ReflectedCheckConstraint, {"name": row[0], "sqltext": row[1]}) for row in rows] @reflection.cache - def get_indexes(self, connection, table_name, schema, **kw): + def get_indexes( + self, connection: Any, table_name: str, schema: Optional[str] = None, **kw: Any + ) -> List[ReflectedIndex]: return [] - # Disable index creation since that's not a thing in Vertica. - # noinspection PyUnusedLocal - def visit_create_index(self, create): - return None + def _get_column_info( + self, + name: str, + format_type: str, + default: Optional[str], + nullable: bool, + schema: Optional[str], + ) -> Dict[str, Any]: + attype = re.sub(r"\(.*\)", "", format_type).strip() + + charlen_match = re.search(r"\(([\d,]+)\)", format_type) + charlen = charlen_match.group(1) if charlen_match else None + + args: Tuple[Any, ...] = () + kwargs: Dict[str, Any] = {} + + if attype in ("numeric", "decimal"): + if charlen: + parts = charlen.split(",") + if len(parts) == 2: + args = (int(parts[0]), int(parts[1])) + elif len(parts) == 1: + args = (int(parts[0]),) + elif attype in ("timestamptz", "timestamp with time zone", "timestamp with timezone"): + if charlen: + args = (int(charlen),) + attype = "timestamptz" + elif attype in ("timetz", "time with time zone", "time with timezone"): + if charlen: + args = (int(charlen),) + attype = "timetz" + elif attype in ("timestamp", "timestamp without time zone"): + if charlen: + args = (int(charlen),) + attype = "timestamp" + elif attype in ("time", "time without time zone"): + if charlen: + args = (int(charlen),) + attype = "time" + elif attype.startswith("interval"): + field_match = re.match(r"interval\s+(.+)", attype, re.I) + if field_match: + kwargs["fields"] = field_match.group(1) + if charlen: + kwargs["precision"] = int(charlen) + attype = "interval" + elif attype in ("geometry", "geography"): + if charlen: + kwargs["srid"] = int(charlen) + elif charlen and "," not in charlen: + args = (int(charlen),) + + coltype_cls = self.ischema_names.get(attype.upper()) + if coltype_cls: + try: + coltype = coltype_cls(*args, **kwargs) + except Exception: + try: + coltype = coltype_cls() + except Exception: + coltype = sqltypes.NULLTYPE + else: + util.warn(f"Did not recognize type '{format_type}' of column '{name}'") + coltype = sqltypes.NULLTYPE + + autoincrement = False + if default is not None: + if "nextval" in default.lower() or "auto_increment" in default.lower() or "identity" in default.lower(): + autoincrement = True + + return { + "name": name, + "type": coltype, + "nullable": nullable, + "default": default, + "autoincrement": autoincrement, + } diff --git a/sqlalchemy_vertica/dialect_pyodbc.py b/sqlalchemy_vertica/dialect_pyodbc.py index 32a9061..50c9398 100644 --- a/sqlalchemy_vertica/dialect_pyodbc.py +++ b/sqlalchemy_vertica/dialect_pyodbc.py @@ -1,10 +1,15 @@ -from __future__ import absolute_import, print_function, division +from __future__ import annotations +from typing import Any from sqlalchemy.connectors.pyodbc import PyODBCConnector from .base import VerticaDialect as BaseVerticaDialect -# noinspection PyAbstractClass, PyClassHasNoInit -class VerticaDialect(PyODBCConnector, BaseVerticaDialect): - pass +class VerticaDialect(PyODBCConnector, BaseVerticaDialect): # type: ignore[misc] + driver = "pyodbc" + supports_statement_cache = True + + @classmethod + def import_dbapi(cls) -> Any: + return PyODBCConnector.import_dbapi() diff --git a/sqlalchemy_vertica/dialect_turbodbc.py b/sqlalchemy_vertica/dialect_turbodbc.py new file mode 100644 index 0000000..4e66c20 --- /dev/null +++ b/sqlalchemy_vertica/dialect_turbodbc.py @@ -0,0 +1,34 @@ +from __future__ import annotations + +from typing import Any, Dict, Sequence, Tuple +from sqlalchemy.engine.url import URL + +from .base import VerticaDialect as BaseVerticaDialect + + +class VerticaDialect(BaseVerticaDialect): + driver = "turbodbc" + supports_statement_cache = True + + @classmethod + def import_dbapi(cls) -> Any: + import turbodbc + return turbodbc + + def create_connect_args(self, url: URL) -> Tuple[Sequence[Any], Dict[str, Any]]: + opts: Dict[str, Any] = {} + if url.host: + opts["host"] = url.host + if url.port is not None: + opts["port"] = int(url.port) + else: + opts["port"] = 5433 + if url.username: + opts["user"] = url.username + if url.password is not None: + opts["password"] = url.password + if url.database: + opts["database"] = url.database + + opts.update(url.query) + return [], opts diff --git a/sqlalchemy_vertica/dialect_vertica_python.py b/sqlalchemy_vertica/dialect_vertica_python.py index 5dfb3e7..6986af0 100644 --- a/sqlalchemy_vertica/dialect_vertica_python.py +++ b/sqlalchemy_vertica/dialect_vertica_python.py @@ -1,13 +1,63 @@ -from __future__ import absolute_import, print_function, division +from __future__ import annotations + +from typing import Any, Dict, Sequence, Tuple +from sqlalchemy.engine.url import URL from .base import VerticaDialect as BaseVerticaDialect -# noinspection PyAbstractClass, PyClassHasNoInit class VerticaDialect(BaseVerticaDialect): - driver = 'vertica_python' + driver = "vertica_python" + supports_statement_cache = True @classmethod - def dbapi(cls): - vertica_python = __import__('vertica_python') + def import_dbapi(cls) -> Any: + import vertica_python return vertica_python + + # Maintain backwards compatibility for any legacy callers + @classmethod + def dbapi(cls) -> Any: # type: ignore[override] + return cls.import_dbapi() + + def create_connect_args(self, url: URL) -> Tuple[Sequence[Any], Dict[str, Any]]: + opts: Dict[str, Any] = {} + + if url.host: + opts["host"] = url.host + if url.port is not None: + try: + opts["port"] = int(url.port) + except (ValueError, TypeError): + opts["port"] = url.port + else: + opts["port"] = 5433 + + if url.username: + opts["user"] = url.username + if url.password is not None: + opts["password"] = url.password + if url.database: + opts["database"] = url.database + + # Query options handling + query = dict(url.query) + + int_keys = {"connection_timeout", "read_timeout", "port"} + bool_keys = {"autocommit", "connection_load_balance"} + + for key, val in query.items(): + if key in int_keys: + try: + opts[key] = int(val) # type: ignore[arg-type] + except (ValueError, TypeError): + opts[key] = val + elif key in bool_keys: + if isinstance(val, str): + opts[key] = val.lower() in ("true", "1", "yes", "on") + else: + opts[key] = bool(val) + else: + opts[key] = val + + return [], opts diff --git a/sqlalchemy_vertica/dialect_vertica_python_async.py b/sqlalchemy_vertica/dialect_vertica_python_async.py new file mode 100644 index 0000000..2169e93 --- /dev/null +++ b/sqlalchemy_vertica/dialect_vertica_python_async.py @@ -0,0 +1,156 @@ +from __future__ import annotations + +import asyncio +from typing import Any, List, Optional, Sequence +from sqlalchemy import pool +from sqlalchemy.connectors.asyncio import ( + AsyncAdapt_dbapi_connection, + AsyncAdapt_dbapi_module, + await_only, +) +from sqlalchemy.engine.url import URL + +from .dialect_vertica_python import VerticaDialect as VerticaDialect_sync + + +class AsyncVerticaCursor: + __slots__ = ("_sync_cursor",) + + def __init__(self, sync_cursor: Any) -> None: + self._sync_cursor = sync_cursor + + async def __aenter__(self) -> AsyncVerticaCursor: + return self + + async def __aexit__(self, exc_type: Any, exc_val: Any, exc_tb: Any) -> None: + await self.close() + + @property + def description(self) -> Any: + return self._sync_cursor.description + + @property + def rowcount(self) -> int: + return getattr(self._sync_cursor, "rowcount", -1) + + @property + def arraysize(self) -> int: + return getattr(self._sync_cursor, "arraysize", 1) + + @arraysize.setter + def arraysize(self, value: int) -> None: + self._sync_cursor.arraysize = value + + async def execute(self, operation: Any, parameters: Optional[Any] = None) -> Any: + if parameters is not None: + return await asyncio.to_thread(self._sync_cursor.execute, operation, parameters) + return await asyncio.to_thread(self._sync_cursor.execute, operation) + + async def executemany(self, operation: Any, seq_of_parameters: Sequence[Any]) -> Any: + return await asyncio.to_thread(self._sync_cursor.executemany, operation, seq_of_parameters) + + async def fetchone(self) -> Optional[Any]: + return await asyncio.to_thread(self._sync_cursor.fetchone) + + async def fetchmany(self, size: Optional[int] = None) -> List[Any]: + if size is not None: + return await asyncio.to_thread(self._sync_cursor.fetchmany, size) + return await asyncio.to_thread(self._sync_cursor.fetchmany) + + async def fetchall(self) -> List[Any]: + return await asyncio.to_thread(self._sync_cursor.fetchall) + + async def close(self) -> None: + try: + await asyncio.to_thread(self._sync_cursor.close) + except Exception: + pass + + def __getattr__(self, name: str) -> Any: + return getattr(self._sync_cursor, name) + + def __setattr__(self, name: str, value: Any) -> None: + if name == "_sync_cursor": + super().__setattr__(name, value) + else: + setattr(self._sync_cursor, name, value) + + +class AsyncVerticaConnection: + __slots__ = ("_sync_conn",) + + def __init__(self, sync_conn: Any) -> None: + self._sync_conn = sync_conn + + def cursor(self, *args: Any, **kwargs: Any) -> Any: + sync_cur = self._sync_conn.cursor(*args, **kwargs) + return AsyncVerticaCursor(sync_cur) + + async def commit(self) -> None: + await asyncio.to_thread(self._sync_conn.commit) + + async def rollback(self) -> None: + await asyncio.to_thread(self._sync_conn.rollback) + + async def close(self) -> None: + try: + await asyncio.to_thread(self._sync_conn.close) + except Exception: + pass + + def __getattr__(self, name: str) -> Any: + return getattr(self._sync_conn, name) + + def __setattr__(self, name: str, value: Any) -> None: + if name == "_sync_conn": + super().__setattr__(name, value) + else: + setattr(self._sync_conn, name, value) + + +class AsyncAdapt_vertica_python_dbapi(AsyncAdapt_dbapi_module): + def __init__(self, vertica_python_module: Any) -> None: + self.vertica_python = vertica_python_module + self.paramstyle = "named" + self._init_dbapi_attributes() + + def _init_dbapi_attributes(self) -> None: + for name in ( + "DatabaseError", + "DataError", + "Error", + "IntegrityError", + "InterfaceError", + "InternalError", + "NotSupportedError", + "OperationalError", + "ProgrammingError", + "Warning", + ): + setattr(self, name, getattr(self.vertica_python, name, Exception)) + + async def async_connect(self, *args: Any, **kwargs: Any) -> AsyncVerticaConnection: + sync_conn = await asyncio.to_thread(self.vertica_python.connect, *args, **kwargs) + return AsyncVerticaConnection(sync_conn) + + def connect(self, *args: Any, **kwargs: Any) -> AsyncAdapt_dbapi_connection: + return AsyncAdapt_dbapi_connection( + self, + await_only(self.async_connect(*args, **kwargs)), + ) + + +class VerticaDialect_vertica_python_async(VerticaDialect_sync): + driver = "vertica_python_async" + is_async = True + supports_statement_cache = True + has_terminate = False + + @classmethod + def import_dbapi(cls) -> Any: + import vertica_python + return AsyncAdapt_vertica_python_dbapi(vertica_python) + + @classmethod + def get_pool_class(cls, url: URL) -> type[pool.Pool]: + return pool.AsyncAdaptedQueuePool diff --git a/sqlalchemy_vertica/types.py b/sqlalchemy_vertica/types.py new file mode 100644 index 0000000..3d931c1 --- /dev/null +++ b/sqlalchemy_vertica/types.py @@ -0,0 +1,215 @@ +from __future__ import annotations + +from typing import Any, Optional, Type +from sqlalchemy import types as sqltypes +from sqlalchemy.types import ( + BIGINT, + BLOB, + BOOLEAN, + CHAR, + DATE, + DATETIME, + DECIMAL, + FLOAT, + INTEGER, + NUMERIC, + REAL, + SMALLINT, + VARCHAR, +) + + +class LONG_VARCHAR(sqltypes.String): + __visit_name__ = "LONG_VARCHAR" + + def __init__(self, length: Optional[int] = None, **kwargs: Any) -> None: + super().__init__(length=length, **kwargs) + + +class LONG_VARBINARY(sqltypes.LargeBinary): + __visit_name__ = "LONG_VARBINARY" + + def __init__(self, length: Optional[int] = None, **kwargs: Any) -> None: + super().__init__(length=length, **kwargs) + + +class VARBINARY(sqltypes.LargeBinary): + __visit_name__ = "VARBINARY" + + def __init__(self, length: Optional[int] = None, **kwargs: Any) -> None: + super().__init__(length=length, **kwargs) + + +class BYTEA(sqltypes.LargeBinary): + __visit_name__ = "BYTEA" + + def __init__(self, length: Optional[int] = None, **kwargs: Any) -> None: + super().__init__(length=length, **kwargs) + + +class RAW(sqltypes.LargeBinary): + __visit_name__ = "RAW" + + def __init__(self, length: Optional[int] = None, **kwargs: Any) -> None: + super().__init__(length=length, **kwargs) + + +class DOUBLE_PRECISION(sqltypes.Float): + __visit_name__ = "DOUBLE_PRECISION" + + def __init__(self, precision: Optional[int] = None, **kwargs: Any) -> None: + super().__init__(precision=precision, **kwargs) + + +class TIME(sqltypes.TIME): + __visit_name__ = "TIME" + + def __init__(self, precision: Optional[int] = None, timezone: bool = False, **kwargs: Any) -> None: + self.precision = precision + super().__init__(timezone=timezone, **kwargs) + + +class TIMETZ(sqltypes.TIME): + __visit_name__ = "TIMETZ" + + def __init__(self, precision: Optional[int] = None, **kwargs: Any) -> None: + self.precision = precision + kwargs.pop("timezone", None) + super().__init__(timezone=True, **kwargs) + + +class TIMESTAMP(sqltypes.TIMESTAMP): + __visit_name__ = "TIMESTAMP" + + def __init__(self, precision: Optional[int] = None, timezone: bool = False, **kwargs: Any) -> None: + self.precision = precision + super().__init__(timezone=timezone, **kwargs) + + +class TIMESTAMPTZ(sqltypes.TIMESTAMP): + __visit_name__ = "TIMESTAMPTZ" + + def __init__(self, precision: Optional[int] = None, **kwargs: Any) -> None: + self.precision = precision + kwargs.pop("timezone", None) + super().__init__(timezone=True, **kwargs) + + +class GEOMETRY(sqltypes.UserDefinedType): + __visit_name__ = "GEOMETRY" + + def __init__(self, srid: Optional[int] = None) -> None: + self.srid = srid + + def get_col_spec(self, **kw: Any) -> str: + if self.srid is not None: + return f"GEOMETRY({self.srid})" + return "GEOMETRY" + + +class GEOGRAPHY(sqltypes.UserDefinedType): + __visit_name__ = "GEOGRAPHY" + + def __init__(self, srid: Optional[int] = None) -> None: + self.srid = srid + + def get_col_spec(self, **kw: Any) -> str: + if self.srid is not None: + return f"GEOGRAPHY({self.srid})" + return "GEOGRAPHY" + + +class UUID(sqltypes.UUID): + __visit_name__ = "UUID" + + +class INTERVAL(sqltypes.TypeEngine): + __visit_name__ = "INTERVAL" + + def __init__( + self, + fields: Optional[str] = None, + precision: Optional[int] = None, + ) -> None: + self.fields = fields + self.precision = precision + + +class ARRAY(sqltypes.TypeEngine): + __visit_name__ = "ARRAY" + + def __init__( + self, + item_type: Type[sqltypes.TypeEngine] | sqltypes.TypeEngine, + length: Optional[int] = None, + ) -> None: + if isinstance(item_type, type): + self.item_type = item_type() + else: + self.item_type = item_type + self.length = length + + +class MAP(sqltypes.TypeEngine): + __visit_name__ = "MAP" + + def __init__( + self, + key_type: Type[sqltypes.TypeEngine] | sqltypes.TypeEngine, + value_type: Type[sqltypes.TypeEngine] | sqltypes.TypeEngine, + ) -> None: + if isinstance(key_type, type): + self.key_type = key_type() + else: + self.key_type = key_type + + if isinstance(value_type, type): + self.value_type = value_type() + else: + self.value_type = value_type + + +class ROW(sqltypes.TypeEngine): + __visit_name__ = "ROW" + + def __init__(self, **fields: Type[sqltypes.TypeEngine] | sqltypes.TypeEngine) -> None: + self.fields = { + k: v() if isinstance(v, type) else v for k, v in fields.items() + } + + +BINARY = VARBINARY + +__all__ = [ + "INTEGER", + "BIGINT", + "SMALLINT", + "CHAR", + "VARCHAR", + "LONG_VARCHAR", + "NUMERIC", + "DECIMAL", + "FLOAT", + "REAL", + "DOUBLE_PRECISION", + "BOOLEAN", + "DATE", + "TIME", + "TIMETZ", + "TIMESTAMP", + "TIMESTAMPTZ", + "DATETIME", + "INTERVAL", + "BLOB", + "BYTEA", + "RAW", + "BINARY", + "VARBINARY", + "LONG_VARBINARY", + "UUID", + "GEOMETRY", + "GEOGRAPHY", + "ARRAY", + "MAP", + "ROW", +] diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..9fe52d1 --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,113 @@ +from __future__ import annotations + +from typing import Any, List, Optional, Tuple +import pytest +from sqlalchemy import Column, ForeignKey, Integer, MetaData, String, Table + + +class MockCursor: + def __init__(self, query_handler: Optional[Any] = None) -> None: + self.query_handler = query_handler + self.description: Optional[List[Tuple[Any, ...]]] = None + self.rowcount: int = -1 + self.arraysize: int = 1 + self._rows: List[Tuple[Any, ...]] = [] + self._closed: bool = False + + def execute(self, operation: Any, parameters: Optional[Any] = None) -> MockCursor: + if self.query_handler: + self.description, self._rows = self.query_handler(operation, parameters) + self.rowcount = len(self._rows) + else: + self.description = [("col1", 1, None, None, None, None, None)] + self._rows = [(1,)] + self.rowcount = 1 + return self + + def executemany(self, operation: Any, seq_of_parameters: Any) -> MockCursor: + self.rowcount = len(seq_of_parameters) + return self + + def fetchone(self) -> Optional[Tuple[Any, ...]]: + if self._rows: + return self._rows.pop(0) + return None + + def fetchmany(self, size: Optional[int] = None) -> List[Tuple[Any, ...]]: + if size is None: + size = self.arraysize + res = self._rows[:size] + self._rows = self._rows[size:] + return res + + def fetchall(self) -> List[Tuple[Any, ...]]: + res = list(self._rows) + self._rows = [] + return res + + def close(self) -> None: + self._closed = True + + +class MockConnection: + def __init__(self, query_handler: Optional[Any] = None) -> None: + self.query_handler = query_handler + self.committed: bool = False + self.rolled_back: bool = False + self.closed: bool = False + + def cursor(self, *args: Any, **kwargs: Any) -> MockCursor: + return MockCursor(self.query_handler) + + def commit(self) -> None: + self.committed = True + + def rollback(self) -> None: + self.rolled_back = True + + def close(self) -> None: + self.closed = True + + +class MockDBAPI: + def __init__(self, query_handler: Optional[Any] = None) -> None: + self.query_handler = query_handler + self.paramstyle = "named" + self.Error = Exception + self.DatabaseError = Exception + self.DataError = Exception + self.IntegrityError = Exception + self.InterfaceError = Exception + self.InternalError = Exception + self.NotSupportedError = Exception + self.OperationalError = Exception + self.ProgrammingError = Exception + self.Warning = Exception + + def connect(self, *args: Any, **kwargs: Any) -> MockConnection: + return MockConnection(self.query_handler) + + +@pytest.fixture +def mock_dbapi() -> MockDBAPI: + return MockDBAPI() + + +@pytest.fixture +def sample_metadata() -> MetaData: + metadata = MetaData() + Table( + "users", + metadata, + Column("id", Integer, primary_key=True, autoincrement=True), + Column("name", String(50), nullable=False), + Column("email", String(100), unique=True), + ) + Table( + "addresses", + metadata, + Column("id", Integer, primary_key=True), + Column("user_id", Integer, ForeignKey("users.id")), + Column("address", String(200)), + ) + return metadata diff --git a/tests/test_alembic.py b/tests/test_alembic.py new file mode 100644 index 0000000..3d9ac55 --- /dev/null +++ b/tests/test_alembic.py @@ -0,0 +1,96 @@ +from __future__ import annotations + +import logging +from unittest.mock import MagicMock +import pytest +from sqlalchemy import Column, Index, Integer, MetaData, String, Table + +from sqlalchemy_vertica.alembic import VerticaImpl +from sqlalchemy_vertica.base import VerticaDialect + + +def test_alembic_vertica_impl_registered() -> None: + try: + from alembic.ddl.impl import DefaultImpl + impl_cls = DefaultImpl.get_by_dialect(VerticaDialect()) + assert issubclass(impl_cls, VerticaImpl) + assert impl_cls.transactional_ddl is True + except ImportError: + pytest.skip("Alembic not installed") + + +def test_alembic_index_operations_are_noops(caplog: pytest.LogCaptureFixture) -> None: + if VerticaImpl is None: + pytest.skip("Alembic not installed") + + dialect = VerticaDialect() + impl = VerticaImpl( + dialect=dialect, + connection=MagicMock(), + as_sql=False, + transactional_ddl=True, + output_buffer=None, + context_opts={}, + ) + + meta = MetaData() + tbl = Table("test_table", meta, Column("id", Integer), Column("name", String(50))) + idx = Index("ix_test_name", tbl.c.name) + + with caplog.at_level(logging.WARNING): + impl.create_index(idx) + assert "Ignoring create_index for 'ix_test_name'" in caplog.text + + caplog.clear() + with caplog.at_level(logging.WARNING): + impl.drop_index(idx) + assert "Ignoring drop_index for 'ix_test_name'" in caplog.text + + +def test_alembic_type_synonyms() -> None: + if VerticaImpl is None: + pytest.skip("Alembic not installed") + + synonyms = VerticaImpl.type_synonyms + + # Check integer synonyms + assert any({"INT", "INTEGER", "INT8", "BIGINT"}.issubset(s) for s in synonyms) + + # Check float synonyms + assert any({"FLOAT", "FLOAT8", "DOUBLE PRECISION", "REAL"}.issubset(s) for s in synonyms) + + # Check varchar synonyms + assert any({"VARCHAR", "VARCHAR2", "TEXT", "LONG VARCHAR"}.issubset(s) for s in synonyms) + + # Check timestamp synonyms + assert any({"TIMESTAMP", "TIMESTAMP WITHOUT TIME ZONE", "DATETIME", "SMALLDATETIME"}.issubset(s) for s in synonyms) + + +def test_alembic_migration_operations_simulation() -> None: + try: + from alembic.migration import MigrationContext + from alembic.operations import Operations + except ImportError: + pytest.skip("Alembic not installed") + + mock_conn = MagicMock() + mock_conn.dialect = VerticaDialect() + + ctx = MigrationContext.configure(mock_conn) + op = Operations(ctx) + + # Test creating a table with op + op.create_table( + "users", + Column("id", Integer, primary_key=True), + Column("username", String(50)), + ) + + # Test adding a column + op.add_column("users", Column("email", String(100))) + + # Test dropping a column + op.drop_column("users", "email") + + # Test creating an index (should be a safe no-op on Vertica) + op.create_index("ix_users_username", "users", ["username"]) diff --git a/tests/test_async.py b/tests/test_async.py new file mode 100644 index 0000000..a70ab9d --- /dev/null +++ b/tests/test_async.py @@ -0,0 +1,191 @@ +from __future__ import annotations + +from typing import Any, List, Optional, Tuple +from unittest.mock import MagicMock, patch +import pytest +from sqlalchemy import pool, text +from sqlalchemy.engine.url import make_url +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine + +from sqlalchemy_vertica.dialect_vertica_python_async import ( + AsyncAdapt_vertica_python_dbapi, + AsyncVerticaConnection, + AsyncVerticaCursor, + VerticaDialect_vertica_python_async, +) + + +class MockSyncCursorForAsync: + def __init__(self) -> None: + self.description = [("val", 1, None, None, None, None, None)] + self.rowcount = 1 + self.arraysize = 1 + self.closed = False + self._rows = [(100,), (200,), (300,)] + + def execute(self, operation: Any, parameters: Optional[Any] = None) -> None: + pass + + def executemany(self, operation: Any, seq_of_parameters: Any) -> None: + self.rowcount = len(seq_of_parameters) + + def fetchone(self) -> Optional[Tuple[Any, ...]]: + if self._rows: + return self._rows.pop(0) + return None + + def fetchmany(self, size: Optional[int] = None) -> List[Tuple[Any, ...]]: + if size is None: + size = self.arraysize + res = self._rows[:size] + self._rows = self._rows[size:] + return res + + def fetchall(self) -> List[Tuple[Any, ...]]: + res = list(self._rows) + self._rows = [] + return res + + def close(self) -> None: + self.closed = True + + +class FailingCursor(MockSyncCursorForAsync): + def close(self) -> None: + raise RuntimeError("Close error") + + +class MockSyncConnForAsync: + def __init__(self) -> None: + self.committed = False + self.rolled_back = False + self.closed = False + + def cursor(self, *args: Any, **kwargs: Any) -> MockSyncCursorForAsync: + return MockSyncCursorForAsync() + + def commit(self) -> None: + self.committed = True + + def rollback(self) -> None: + self.rolled_back = True + + def close(self) -> None: + self.closed = True + + +class FailingConn(MockSyncConnForAsync): + def close(self) -> None: + raise RuntimeError("Close conn error") + + +@pytest.mark.asyncio +async def test_async_cursor_methods() -> None: + sync_cur = MockSyncCursorForAsync() + async_cur = AsyncVerticaCursor(sync_cur) + + assert async_cur.description == sync_cur.description + assert async_cur.rowcount == 1 + assert async_cur.arraysize == 1 + + async_cur.arraysize = 2 + assert async_cur.arraysize == 2 + + # Context manager test + async with async_cur as cur: + await cur.execute("SELECT 1") + await cur.executemany("INSERT INTO t VALUES (?)", [(1,), (2,)]) + row1 = await cur.fetchone() + assert row1 == (100,) + rows = await cur.fetchmany() + assert rows == [(200,), (300,)] + all_rows = await cur.fetchall() + assert all_rows == [] + + assert sync_cur.closed is True + + +@pytest.mark.asyncio +async def test_async_cursor_close_exception() -> None: + failing_cur = FailingCursor() + async_cur = AsyncVerticaCursor(failing_cur) + # Should not raise exception + await async_cur.close() + + +@pytest.mark.asyncio +async def test_async_connection_methods() -> None: + sync_conn = MockSyncConnForAsync() + async_conn = AsyncVerticaConnection(sync_conn) + + cur = async_conn.cursor() + assert isinstance(cur, AsyncVerticaCursor) + + await async_conn.commit() + assert sync_conn.committed is True + + await async_conn.rollback() + assert sync_conn.rolled_back is True + + await async_conn.close() + assert sync_conn.closed is True + + +@pytest.mark.asyncio +async def test_async_connection_close_exception() -> None: + failing_conn = FailingConn() + async_conn = AsyncVerticaConnection(failing_conn) + # Should not raise exception + await async_conn.close() + + +def test_async_dialect_pool_and_import() -> None: + url = make_url("vertica+vertica_python_async://") + pool_cls = VerticaDialect_vertica_python_async.get_pool_class(url) + assert pool_cls is pool.AsyncAdaptedQueuePool + + mock_vp = MagicMock() + with patch.dict("sys.modules", {"vertica_python": mock_vp}): + dbapi = VerticaDialect_vertica_python_async.import_dbapi() + assert isinstance(dbapi, AsyncAdapt_vertica_python_dbapi) + assert hasattr(dbapi, "OperationalError") + + +@pytest.mark.asyncio +async def test_create_async_engine_execution() -> None: + mock_vertica = MagicMock() + mock_vertica.connect.return_value = MockSyncConnForAsync() + mock_vertica.Error = Exception + mock_vertica.DatabaseError = Exception + mock_vertica.OperationalError = Exception + mock_vertica.IntegrityError = Exception + mock_vertica.ProgrammingError = Exception + + mock_dbapi = AsyncAdapt_vertica_python_dbapi(mock_vertica) + + engine = create_async_engine( + "vertica+vertica_python_async://user:pass@localhost:5433/testdb", + module=mock_dbapi, + ) + + # 1. Test connect and query + async with engine.connect() as conn: + res = await conn.execute(text("SELECT 100")) + row = res.fetchone() + assert row is not None + assert row[0] == 100 + + # 2. Test transaction begin block + async with engine.begin() as conn: + res = await conn.execute(text("SELECT 100")) + val = res.scalar() + assert val == 100 + + # 3. Test AsyncSession + async_session = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False) + async with async_session() as session: + sess_res = await session.execute(text("SELECT 100")) + val = sess_res.scalar() + assert val == 100 + + await engine.dispose() diff --git a/tests/test_compiler.py b/tests/test_compiler.py new file mode 100644 index 0000000..4552bb5 --- /dev/null +++ b/tests/test_compiler.py @@ -0,0 +1,100 @@ +from __future__ import annotations + +import pytest +from sqlalchemy import ( + CheckConstraint, + Column, + ForeignKey, + Index, + Integer, + MetaData, + Sequence, + String, + Table, + UniqueConstraint, + select, +) +from sqlalchemy.schema import CreateIndex, CreateTable, DropIndex + +from sqlalchemy_vertica.base import VerticaDialect + + +@pytest.fixture +def dialect() -> VerticaDialect: + return VerticaDialect() + + +def test_create_table_autoincrement(dialect: VerticaDialect) -> None: + meta = MetaData() + tbl = Table( + "users", + meta, + Column("id", Integer, primary_key=True, autoincrement=True), + Column("name", String(50), nullable=False), + ) + stmt = str(CreateTable(tbl).compile(dialect=dialect)).strip() + assert "CREATE TABLE users" in stmt + assert "id INT AUTO_INCREMENT NOT NULL" in stmt + assert "name VARCHAR(50) NOT NULL" in stmt + assert "PRIMARY KEY (id)" in stmt + + +def test_create_table_with_constraints(dialect: VerticaDialect) -> None: + meta = MetaData() + Table( + "parents", + meta, + Column("id", Integer, primary_key=True), + ) + child = Table( + "children", + meta, + Column("id", Integer, primary_key=True), + Column("parent_id", Integer, ForeignKey("parents.id", name="fk_child_parent")), + Column("code", String(20), nullable=False), + Column("age", Integer), + UniqueConstraint("code", name="uq_child_code"), + CheckConstraint("age >= 0", name="chk_age_pos"), + ) + + stmt = str(CreateTable(child).compile(dialect=dialect)).strip() + assert "CREATE TABLE children" in stmt + assert "CONSTRAINT fk_child_parent FOREIGN KEY (parent_id) REFERENCES parents (id)" in stmt + assert "CONSTRAINT uq_child_code UNIQUE (code)" in stmt + assert "CONSTRAINT chk_age_pos CHECK (age >= 0)" in stmt + + +def test_index_compilation_is_noop(dialect: VerticaDialect) -> None: + meta = MetaData() + tbl = Table("orders", meta, Column("id", Integer, primary_key=True), Column("code", String(20))) + idx = Index("ix_orders_code", tbl.c.code) + + create_idx_sql = str(CreateIndex(idx).compile(dialect=dialect)).strip() + assert create_idx_sql == "" + + drop_idx_sql = str(DropIndex(idx).compile(dialect=dialect)).strip() + assert drop_idx_sql == "" + + +def test_limit_offset_clause(dialect: VerticaDialect) -> None: + meta = MetaData() + tbl = Table("items", meta, Column("id", Integer, primary_key=True), Column("name", String(50))) + + stmt = select(tbl).limit(10).offset(20) + sql_text = str(stmt.compile(dialect=dialect, compile_kwargs={"literal_binds": True})) + assert "LIMIT 10" in sql_text + assert "OFFSET 20" in sql_text + + +def test_sequence_nextval(dialect: VerticaDialect) -> None: + seq = Sequence("my_vertica_seq") + sql_text = str(select(seq.next_value()).compile(dialect=dialect)) + assert "my_vertica_seq.NEXTVAL" in sql_text + + +def test_for_update_clause(dialect: VerticaDialect) -> None: + meta = MetaData() + tbl = Table("accounts", meta, Column("id", Integer, primary_key=True), Column("balance", Integer)) + stmt = select(tbl).with_for_update() + sql_text = str(stmt.compile(dialect=dialect)) + assert "FOR UPDATE" in sql_text diff --git a/tests/test_dialect.py b/tests/test_dialect.py new file mode 100644 index 0000000..366c750 --- /dev/null +++ b/tests/test_dialect.py @@ -0,0 +1,142 @@ +from __future__ import annotations + +from typing import Any, Optional +import sqlalchemy as sa +from sqlalchemy.engine.url import make_url + +from sqlalchemy_vertica.base import VerticaDialect +from sqlalchemy_vertica.dialect_vertica_python import VerticaDialect as VerticaPythonDialect +from sqlalchemy_vertica.dialect_vertica_python_async import ( + VerticaDialect_vertica_python_async as VerticaPythonAsyncDialect, +) +from sqlalchemy_vertica.dialect_pyodbc import VerticaDialect as PyODBCDialect +from sqlalchemy_vertica.dialect_turbodbc import VerticaDialect as TurbodbcDialect + + +def test_dialect_registry() -> None: + d_sync = sa.dialects.registry.load("vertica") + assert issubclass(d_sync, VerticaDialect) + + d_vp = sa.dialects.registry.load("vertica.vertica_python") + assert issubclass(d_vp, VerticaPythonDialect) + + d_vp_async = sa.dialects.registry.load("vertica.vertica_python_async") + assert issubclass(d_vp_async, VerticaPythonAsyncDialect) + assert d_vp_async.is_async is True + + d_async_vp = sa.dialects.registry.load("vertica.async_vertica_python") + assert issubclass(d_async_vp, VerticaPythonAsyncDialect) + + d_pyodbc = sa.dialects.registry.load("vertica.pyodbc") + assert issubclass(d_pyodbc, PyODBCDialect) + + d_turbodbc = sa.dialects.registry.load("vertica.turbodbc") + assert issubclass(d_turbodbc, TurbodbcDialect) + + +def test_statement_cache_enabled() -> None: + assert VerticaDialect.supports_statement_cache is True + assert VerticaPythonDialect.supports_statement_cache is True + assert VerticaPythonAsyncDialect.supports_statement_cache is True + assert PyODBCDialect.supports_statement_cache is True + assert TurbodbcDialect.supports_statement_cache is True + + +def test_create_connect_args_sync() -> None: + dialect = VerticaPythonDialect() + url = make_url( + "vertica+vertica_python://myuser:mypass@dbhost:5433/mydb" + "?connection_timeout=15&read_timeout=30&autocommit=true" + "&connection_load_balance=1&unicode_error=strict&session_label=app1" + ) + cargs, cparams = dialect.create_connect_args(url) + + assert cargs == [] + assert cparams["host"] == "dbhost" + assert cparams["port"] == 5433 + assert cparams["user"] == "myuser" + assert cparams["password"] == "mypass" + assert cparams["database"] == "mydb" + assert cparams["connection_timeout"] == 15 + assert cparams["read_timeout"] == 30 + assert cparams["autocommit"] is True + assert cparams["connection_load_balance"] is True + assert cparams["unicode_error"] == "strict" + assert cparams["session_label"] == "app1" + + +def test_create_connect_args_invalid_int_and_bool() -> None: + dialect = VerticaPythonDialect() + url = make_url("vertica+vertica_python://dbhost/mydb?port=invalid&autocommit=false&connection_load_balance=False") + _, cparams = dialect.create_connect_args(url) + assert cparams["port"] == "invalid" + assert cparams["autocommit"] is False + assert cparams["connection_load_balance"] is False + + +def test_create_connect_args_default_port() -> None: + dialect = VerticaPythonDialect() + url = make_url("vertica+vertica_python://myuser:mypass@dbhost/mydb") + _, cparams = dialect.create_connect_args(url) + assert cparams["port"] == 5433 + + +def test_base_dialect_create_connect_args() -> None: + dialect = VerticaDialect() + url = make_url("vertica://myuser:mypass@dbhost:5433/mydb?backup_server_node=node2") + cargs, cparams = dialect.create_connect_args(url) + assert cargs == [] + assert cparams["user"] == "myuser" + assert cparams["backup_server_node"] == "node2" + + +def test_server_version_info_parsing() -> None: + dialect = VerticaDialect() + + class MockConn: + def __init__(self, ver_str: Optional[str]) -> None: + self.ver_str = ver_str + + def scalar(self, stmt: Any, params: Any = None) -> Optional[str]: + return self.ver_str + + # Test Vertica 24.1 + conn = MockConn("Vertica Analytic Database v24.1.0-0") + assert dialect._get_server_version_info(conn) == (24, 1, 0) + + # Test Vertica 12.0.4 + conn = MockConn("Vertica Analytic Database v12.0.4-1") + assert dialect._get_server_version_info(conn) == (12, 0, 4) + + # Test OpenText Vertica 23.4 + conn = MockConn("OpenText Vertica Analytic Database v23.4.0") + assert dialect._get_server_version_info(conn) == (23, 4, 0) + + # Test Vertica 11.1 + conn = MockConn("Vertica Analytic Database v11.1.1-0") + assert dialect._get_server_version_info(conn) == (11, 1, 1) + + # Test unknown banner fallback + conn = MockConn("Custom DB 1.0") + assert dialect._get_server_version_info(conn) == (0, 0, 0) + + # Test None version + conn = MockConn(None) + assert dialect._get_server_version_info(conn) == (0, 0, 0) + + +def test_default_schema_name() -> None: + dialect = VerticaDialect() + + class MockConn: + def scalar(self, stmt: Any, params: Any = None) -> Optional[str]: + return "analytics" + + conn = MockConn() + assert dialect._get_default_schema_name(conn) == "analytics" + + class MockConnNone: + def scalar(self, stmt: Any, params: Any = None) -> Optional[str]: + return None + + assert dialect._get_default_schema_name(MockConnNone()) == "public" diff --git a/tests/test_drivers.py b/tests/test_drivers.py new file mode 100644 index 0000000..12b937c --- /dev/null +++ b/tests/test_drivers.py @@ -0,0 +1,59 @@ +from __future__ import annotations + +from unittest.mock import MagicMock, patch +import pytest +from sqlalchemy.engine.url import make_url + +from sqlalchemy_vertica.dialect_vertica_python import VerticaDialect as VerticaPythonDialect +from sqlalchemy_vertica.dialect_pyodbc import VerticaDialect as PyODBCDialect +from sqlalchemy_vertica.dialect_turbodbc import VerticaDialect as TurbodbcDialect + + +def test_vertica_python_driver_import_dbapi() -> None: + try: + dbapi = VerticaPythonDialect.import_dbapi() + assert dbapi is not None + assert VerticaPythonDialect.dbapi() is dbapi + except ImportError: + pytest.skip("vertica-python not installed") + + +def test_pyodbc_driver_connect_args() -> None: + dialect = PyODBCDialect() + url = make_url("vertica+pyodbc:///?odbc_connect=DSN%3DVerticaDSN") + cargs, cparams = dialect.create_connect_args(url) + assert cargs == ["DSN=VerticaDSN"] or "DSN=VerticaDSN" in str(cargs) or "DSN=VerticaDSN" in str(cparams) + + +def test_pyodbc_driver_import_dbapi() -> None: + mock_pyodbc = MagicMock() + with patch("sqlalchemy.connectors.pyodbc.PyODBCConnector.import_dbapi", return_value=mock_pyodbc): + dbapi = PyODBCDialect.import_dbapi() + assert dbapi is mock_pyodbc + + +def test_turbodbc_driver_connect_args() -> None: + dialect = TurbodbcDialect() + url = make_url("vertica+turbodbc://myuser:mypass@dbhost:5433/mydb?read_buffer_size=5000") + cargs, cparams = dialect.create_connect_args(url) + assert cargs == [] + assert cparams["host"] == "dbhost" + assert cparams["port"] == 5433 + assert cparams["user"] == "myuser" + assert cparams["password"] == "mypass" + assert cparams["database"] == "mydb" + assert cparams["read_buffer_size"] == "5000" + + +def test_turbodbc_driver_connect_args_defaults() -> None: + dialect = TurbodbcDialect() + url = make_url("vertica+turbodbc://") + cargs, cparams = dialect.create_connect_args(url) + assert cparams["port"] == 5433 + + +def test_turbodbc_driver_import_dbapi() -> None: + mock_turbodbc = MagicMock() + with patch.dict("sys.modules", {"turbodbc": mock_turbodbc}): + dbapi = TurbodbcDialect.import_dbapi() + assert dbapi is mock_turbodbc diff --git a/tests/test_reflection.py b/tests/test_reflection.py new file mode 100644 index 0000000..fd5c9dc --- /dev/null +++ b/tests/test_reflection.py @@ -0,0 +1,291 @@ +from __future__ import annotations + +from typing import Any, Dict, List, Tuple +import pytest +from sqlalchemy import exc +from sqlalchemy.types import INTEGER, VARCHAR + +from sqlalchemy_vertica.base import VerticaDialect + + +class MockConnectionForReflection: + def __init__(self, routes: Dict[str, Any]) -> None: + self.routes = routes + self.executed_queries: List[Tuple[str, Any]] = [] + + def execute(self, stmt: Any, params: Any = None) -> MockResult: + query = str(stmt).strip() + self.executed_queries.append((query, params)) + + for pattern, res in self.routes.items(): + if pattern in query: + if callable(res): + return MockResult(res(query, params)) + return MockResult(res) + return MockResult([]) + + def scalar(self, stmt: Any, params: Any = None) -> Any: + res = self.execute(stmt, params).fetchall() + if res and res[0]: + return res[0][0] + return None + + +class MockResult: + def __init__(self, rows: List[Tuple[Any, ...]]) -> None: + self._rows = rows + + def fetchall(self) -> List[Tuple[Any, ...]]: + return list(self._rows) + + def scalar(self) -> Any: + if self._rows and self._rows[0]: + return self._rows[0][0] + return None + + def __iter__(self) -> Any: + return iter(self._rows) + + +@pytest.fixture +def dialect() -> VerticaDialect: + return VerticaDialect() + + +def test_get_schema_names(dialect: VerticaDialect) -> None: + conn = MockConnectionForReflection({ + "v_catalog.schemata": [ + ("public",), + ("analytics",), + ("v_catalog",), + ("v_monitor",), + ] + }) + schemas = dialect.get_schema_names(conn) + assert schemas == ["public", "analytics"] + + +def test_get_table_names(dialect: VerticaDialect) -> None: + conn_schema = MockConnectionForReflection({ + "v_catalog.tables": [ + ("users",), + ("orders",), + ] + }) + tables = dialect.get_table_names(conn_schema, schema="public") + assert tables == ["users", "orders"] + + conn_all = MockConnectionForReflection({ + "v_catalog.tables": [ + ("users",), + ("orders",), + ] + }) + tables_all = dialect.get_table_names(conn_all, schema=None) + assert tables_all == ["users", "orders"] + + +def test_get_temp_table_names(dialect: VerticaDialect) -> None: + conn_schema = MockConnectionForReflection({ + "is_temp_table": [ + ("temp_session_data",), + ] + }) + temp_tables = dialect.get_temp_table_names(conn_schema, schema="public") + assert temp_tables == ["temp_session_data"] + + conn_all = MockConnectionForReflection({ + "is_temp_table": [ + ("temp_session_data",), + ] + }) + temp_tables_all = dialect.get_temp_table_names(conn_all, schema=None) + assert temp_tables_all == ["temp_session_data"] + + +def test_get_view_names_and_definition(dialect: VerticaDialect) -> None: + conn_schema = MockConnectionForReflection({ + "v_catalog.views": [ + ("active_users_view",), + ] + }) + views = dialect.get_view_names(conn_schema, schema="public") + assert views == ["active_users_view"] + + conn_all = MockConnectionForReflection({ + "v_catalog.views": [ + ("active_users_view",), + ] + }) + views_all = dialect.get_view_names(conn_all, schema=None) + assert views_all == ["active_users_view"] + + conn_def = MockConnectionForReflection({ + "SELECT current_schema()": [("public",)], + "view_definition": [ + ("SELECT * FROM users WHERE active = true",), + ] + }) + vdef = dialect.get_view_definition(conn_def, "active_users_view", schema=None) + assert vdef == "SELECT * FROM users WHERE active = true" + + +def test_get_pk_constraint(dialect: VerticaDialect) -> None: + conn = MockConnectionForReflection({ + "SELECT current_schema()": [("public",)], + "v_catalog.primary_keys": [ + ("pk_orders", "order_id"), + ("pk_orders", "item_id"), + ] + }) + pk = dialect.get_pk_constraint(conn, "orders", schema=None) + assert pk["name"] == "pk_orders" + assert pk["constrained_columns"] == ["order_id", "item_id"] + + # Test empty PK + conn_empty = MockConnectionForReflection({ + "SELECT current_schema()": [("public",)], + "v_catalog.primary_keys": [] + }) + pk_empty = dialect.get_pk_constraint(conn_empty, "orders", schema="public") + assert pk_empty == {"constrained_columns": [], "name": None} + + +def test_get_foreign_keys(dialect: VerticaDialect) -> None: + conn = MockConnectionForReflection({ + "SELECT current_schema()": [("public",)], + "v_catalog.foreign_keys": [ + ("fk_order_user", "user_id", "public", "users", "id"), + ] + }) + fks = dialect.get_foreign_keys(conn, "orders", schema=None) + assert len(fks) == 1 + assert fks[0]["name"] == "fk_order_user" + assert fks[0]["constrained_columns"] == ["user_id"] + assert fks[0]["referred_schema"] == "public" + assert fks[0]["referred_table"] == "users" + assert fks[0]["referred_columns"] == ["id"] + + # Test empty foreign keys + conn_empty = MockConnectionForReflection({ + "SELECT current_schema()": [("public",)], + "v_catalog.foreign_keys": [] + }) + fks_empty = dialect.get_foreign_keys(conn_empty, "orders", schema="public") + assert fks_empty == [] + + +def test_get_unique_constraints(dialect: VerticaDialect) -> None: + conn = MockConnectionForReflection({ + "SELECT current_schema()": [("public",)], + "constraint_type = 'u'": [ + ("uq_user_email", "email"), + ] + }) + uqs = dialect.get_unique_constraints(conn, "users", schema=None) + assert len(uqs) == 1 + assert uqs[0]["name"] == "uq_user_email" + assert uqs[0]["column_names"] == ["email"] + + +def test_get_check_constraints(dialect: VerticaDialect) -> None: + conn = MockConnectionForReflection({ + "SELECT current_schema()": [("public",)], + "constraint_type = 'c'": [ + ("chk_age", "age >= 18"), + ] + }) + chks = dialect.get_check_constraints(conn, "users", schema=None) + assert len(chks) == 1 + assert chks[0]["name"] == "chk_age" + assert chks[0]["sqltext"] == "age >= 18" + + +def test_get_table_comment(dialect: VerticaDialect) -> None: + conn = MockConnectionForReflection({ + "SELECT current_schema()": [("public",)], + "v_catalog.comments": [ + ("Main user accounts table",), + ] + }) + comment = dialect.get_table_comment(conn, "users", schema=None) + assert comment == {"text": "Main user accounts table"} + + +def test_get_indexes(dialect: VerticaDialect) -> None: + conn = MockConnectionForReflection({}) + assert dialect.get_indexes(conn, "users", schema="public") == [] + + +def test_get_table_oid(dialect: VerticaDialect) -> None: + conn = MockConnectionForReflection({ + "SELECT current_schema()": [("public",)], + "SELECT table_id FROM": [ + (45035996273704976,), + ] + }) + oid = dialect.get_table_oid(conn, "users", schema=None) + assert oid == 45035996273704976 + + # Test not found + conn_empty = MockConnectionForReflection({"SELECT table_id FROM": []}) + with pytest.raises(exc.NoSuchTableError): + dialect.get_table_oid(conn_empty, "nonexistent", schema="public") + + +def test_get_columns(dialect: VerticaDialect) -> None: + conn = MockConnectionForReflection({ + "SELECT current_schema()": [("public",)], + "v_catalog.columns": [ + ("id", "int", "nextval('user_id_seq')", False, 1), + ("username", "varchar(50)", None, False, 2), + ("bio", "varchar(500)", "''", True, 3), + ], + "v_catalog.primary_keys": [ + ("id",), + ], + "v_catalog.comments": [ + ("username", "Unique login handle"), + ], + }) + + cols = dialect.get_columns(conn, "users", schema=None) + assert len(cols) == 3 + + id_col = cols[0] + assert id_col["name"] == "id" + assert isinstance(id_col["type"], INTEGER) + assert id_col["primary_key"] is True # type: ignore[typeddict-item] + assert id_col["autoincrement"] is True + assert id_col["nullable"] is False + + user_col = cols[1] + assert user_col["name"] == "username" + assert isinstance(user_col["type"], VARCHAR) + assert user_col["type"].length == 50 + assert user_col["comment"] == "Unique login handle" + assert user_col["primary_key"] is False # type: ignore[typeddict-item] + + +def test_get_columns_nonexistent_table(dialect: VerticaDialect) -> None: + conn = MockConnectionForReflection({ + "v_catalog.columns": [], + "v_catalog.all_tables": [(False,)], + }) + with pytest.raises(exc.NoSuchTableError): + dialect.get_columns(conn, "missing_table", schema="public") + + +def test_has_table_has_schema_has_sequence(dialect: VerticaDialect) -> None: + conn = MockConnectionForReflection({ + "SELECT current_schema()": [("public",)], + "v_catalog.all_tables": [(True,)], + "v_catalog.schemata": [(True,)], + "v_catalog.sequences": [(True,)], + "v_catalog.types": [(True,)], + }) + + assert dialect.has_table(conn, "users", schema=None) is True + assert dialect.has_schema(conn, "public") is True + assert dialect.has_sequence(conn, "user_seq", schema=None) is True + assert dialect.has_type(conn, "geometry") is True diff --git a/tests/test_types.py b/tests/test_types.py new file mode 100644 index 0000000..85a1cca --- /dev/null +++ b/tests/test_types.py @@ -0,0 +1,183 @@ +from __future__ import annotations + +from typing import Any +import pytest +import sqlalchemy as sa +from sqlalchemy.types import ( + BIGINT, + BLOB, + BOOLEAN, + CHAR, + DATE, + DATETIME, + DECIMAL, + FLOAT, + INTEGER, + NUMERIC, + REAL, + SMALLINT, + VARCHAR, +) + +from sqlalchemy_vertica.base import VerticaDialect +from sqlalchemy_vertica.types import ( + ARRAY, + BYTEA, + DOUBLE_PRECISION, + GEOGRAPHY, + GEOMETRY, + INTERVAL, + LONG_VARBINARY, + LONG_VARCHAR, + MAP, + RAW, + ROW, + TIME, + TIMESTAMPTZ, + TIMETZ, + TIMESTAMP, + UUID, + VARBINARY, +) + + +@pytest.fixture +def dialect() -> VerticaDialect: + return VerticaDialect() + + +def compile_type(dialect: VerticaDialect, type_engine: Any) -> str: + return dialect.type_compiler.process(type_engine) + + +def test_standard_types_compilation(dialect: VerticaDialect) -> None: + assert compile_type(dialect, INTEGER()) == "INT" + assert compile_type(dialect, BIGINT()) == "BIGINT" + assert compile_type(dialect, SMALLINT()) == "SMALLINT" + assert compile_type(dialect, FLOAT()) == "FLOAT" + assert compile_type(dialect, FLOAT(precision=24)) == "FLOAT(24)" + assert compile_type(dialect, DOUBLE_PRECISION()) == "DOUBLE PRECISION" + assert compile_type(dialect, REAL()) == "REAL" + assert compile_type(dialect, NUMERIC()) == "NUMERIC" + assert compile_type(dialect, NUMERIC(10)) == "NUMERIC(10)" + assert compile_type(dialect, NUMERIC(10, 2)) == "NUMERIC(10, 2)" + assert compile_type(dialect, DECIMAL(12, 4)) == "NUMERIC(12, 4)" + assert compile_type(dialect, VARCHAR(100)) == "VARCHAR(100)" + assert compile_type(dialect, VARCHAR()) == "VARCHAR" + assert compile_type(dialect, CHAR(10)) == "CHAR(10)" + assert compile_type(dialect, CHAR()) == "CHAR" + assert compile_type(dialect, sa.Text()) == "LONG VARCHAR" + assert compile_type(dialect, BLOB()) == "LONG VARBINARY" + assert compile_type(dialect, sa.LargeBinary()) == "VARBINARY" + assert compile_type(dialect, sa.LargeBinary(256)) == "VARBINARY(256)" + assert compile_type(dialect, BOOLEAN()) == "BOOLEAN" + assert compile_type(dialect, DATE()) == "DATE" + assert compile_type(dialect, TIME()) == "TIME" + assert compile_type(dialect, TIME(precision=6)) == "TIME(6)" + assert compile_type(dialect, TIME(timezone=True)) == "TIMETZ" + assert compile_type(dialect, TIMESTAMP()) == "TIMESTAMP" + assert compile_type(dialect, TIMESTAMP(precision=3)) == "TIMESTAMP(3)" + assert compile_type(dialect, TIMESTAMP(timezone=True)) == "TIMESTAMPTZ" + assert compile_type(dialect, DATETIME()) == "DATETIME" + + +def test_vertica_specific_types_compilation(dialect: VerticaDialect) -> None: + assert compile_type(dialect, LONG_VARCHAR()) == "LONG VARCHAR" + assert compile_type(dialect, LONG_VARCHAR(length=65000)) == "LONG VARCHAR(65000)" + assert compile_type(dialect, LONG_VARBINARY()) == "LONG VARBINARY" + assert compile_type(dialect, LONG_VARBINARY(length=65000)) == "LONG VARBINARY(65000)" + assert compile_type(dialect, VARBINARY(128)) == "VARBINARY(128)" + assert compile_type(dialect, BYTEA()) == "BYTEA" + assert compile_type(dialect, RAW()) == "RAW" + assert compile_type(dialect, TIMESTAMPTZ()) == "TIMESTAMPTZ" + assert compile_type(dialect, TIMESTAMPTZ(precision=6)) == "TIMESTAMPTZ(6)" + assert compile_type(dialect, TIMETZ()) == "TIMETZ" + assert compile_type(dialect, TIMETZ(precision=3)) == "TIMETZ(3)" + assert compile_type(dialect, INTERVAL(fields="DAY TO SECOND")) == "INTERVAL DAY TO SECOND" + assert compile_type(dialect, INTERVAL(fields="YEAR TO MONTH", precision=2)) == "INTERVAL YEAR TO MONTH (2)" + assert compile_type(dialect, UUID()) == "UUID" + assert compile_type(dialect, GEOMETRY()) == "GEOMETRY" + assert compile_type(dialect, GEOMETRY(srid=4326)) == "GEOMETRY(4326)" + assert GEOMETRY().get_col_spec() == "GEOMETRY" + assert GEOMETRY(srid=4326).get_col_spec() == "GEOMETRY(4326)" + assert compile_type(dialect, GEOGRAPHY()) == "GEOGRAPHY" + assert compile_type(dialect, GEOGRAPHY(srid=4326)) == "GEOGRAPHY(4326)" + assert GEOGRAPHY().get_col_spec() == "GEOGRAPHY" + assert GEOGRAPHY(srid=4326).get_col_spec() == "GEOGRAPHY(4326)" + assert compile_type(dialect, ARRAY(INTEGER)) == "ARRAY[INT]" + assert compile_type(dialect, ARRAY(VARCHAR(50), length=5)) == "ARRAY[VARCHAR(50), 5]" + assert compile_type(dialect, MAP(VARCHAR, INTEGER)) == "MAP[VARCHAR, INT]" + assert compile_type(dialect, ROW(x=INTEGER, y=VARCHAR(50))) in ( + "ROW(x INT, y VARCHAR(50))", + "ROW(y VARCHAR(50), x INT)", + ) + + +def test_column_info_reflection_parsing(dialect: VerticaDialect) -> None: + # Test numeric with only precision + info = dialect._get_column_info("amount", "numeric(18)", None, False, "public") + assert isinstance(info["type"], NUMERIC) + assert info["type"].precision == 18 + + # Test numeric with precision & scale + info = dialect._get_column_info("price", "numeric(10,2)", None, False, "public") + assert isinstance(info["type"], NUMERIC) + assert info["type"].precision == 10 + assert info["type"].scale == 2 + assert info["nullable"] is False + + # Test timestamptz + info = dialect._get_column_info("created_at", "timestamptz(6)", None, True, "public") + assert isinstance(info["type"], TIMESTAMPTZ) + assert info["type"].precision == 6 + assert info["type"].timezone is True + + # Test timestamp without timezone + info = dialect._get_column_info("updated_at", "timestamp(3)", None, True, "public") + assert isinstance(info["type"], TIMESTAMP) + assert info["type"].precision == 3 + assert info["type"].timezone is False + + # Test timetz + info = dialect._get_column_info("event_time", "timetz(3)", None, True, "public") + assert isinstance(info["type"], TIMETZ) + assert info["type"].precision == 3 + assert info["type"].timezone is True + + # Test time without timezone + info = dialect._get_column_info("start_time", "time(2)", None, True, "public") + assert isinstance(info["type"], TIME) + assert info["type"].precision == 2 + assert info["type"].timezone is False + + # Test interval + info = dialect._get_column_info("duration", "interval day to second(6)", None, True, "public") + assert isinstance(info["type"], INTERVAL) + + # Test geometry + info = dialect._get_column_info("geom", "geometry(4326)", None, True, "public") + assert isinstance(info["type"], GEOMETRY) + assert info["type"].srid == 4326 + + # Test geography + info = dialect._get_column_info("geog", "geography(4326)", None, True, "public") + assert isinstance(info["type"], GEOGRAPHY) + assert info["type"].srid == 4326 + + # Test varchar with length + info = dialect._get_column_info("name", "varchar(100)", None, True, "public") + assert isinstance(info["type"], VARCHAR) + assert info["type"].length == 100 + + # Test autoincrement sequence default + info = dialect._get_column_info("id", "int", "nextval('user_id_seq')", False, "public") + assert info["autoincrement"] is True + + # Test identity default + info = dialect._get_column_info("id2", "int", "IDENTITY(1,1)", False, "public") + assert info["autoincrement"] is True + + # Test unknown fallback + with pytest.warns(sa.exc.SAWarning): + info = dialect._get_column_info("custom", "unknown_type_foo", None, True, "public") + assert isinstance(info["type"], sa.sql.sqltypes.NullType)