diff --git a/src/sqlalchemy_declarative_extensions/alembic/base.py b/src/sqlalchemy_declarative_extensions/alembic/base.py index 50469b8..ef00abf 100644 --- a/src/sqlalchemy_declarative_extensions/alembic/base.py +++ b/src/sqlalchemy_declarative_extensions/alembic/base.py @@ -12,10 +12,13 @@ def register_alembic_events( procedures: bool = True, triggers: bool = True, rows: bool = True, + snowflake_dynamic_tables: bool = False, ): """Register handlers into alembic's event system for the supported object types. By default all object types are enabled, but each can be individually disabled. + Snowflake-specific types (e.g. ``snowflake_dynamic_tables``) are opt-in and + default to disabled since they are dialect-specific. Note this is the opposite of the defaults when registering against SQLAlchemy's event system. @@ -47,6 +50,9 @@ def register_alembic_events( if rows: import sqlalchemy_declarative_extensions.alembic.row # noqa + if snowflake_dynamic_tables: + import sqlalchemy_declarative_extensions.alembic.dynamic_table # noqa + def _traverse_any_directive(self, context, revision, directive) -> None: pass diff --git a/src/sqlalchemy_declarative_extensions/alembic/dynamic_table.py b/src/sqlalchemy_declarative_extensions/alembic/dynamic_table.py new file mode 100644 index 0000000..3460ea3 --- /dev/null +++ b/src/sqlalchemy_declarative_extensions/alembic/dynamic_table.py @@ -0,0 +1,41 @@ +from __future__ import annotations + +from alembic.autogenerate.api import AutogenContext + +from sqlalchemy_declarative_extensions.alembic.base import ( + register_comparator_dispatcher, + register_renderer_dispatcher, + register_rewriter_dispatcher, +) +from sqlalchemy_declarative_extensions.dialects.snowflake.dynamic_table import ( + CreateDynamicTableOp, + DropDynamicTableOp, + DynamicTableOperation, + DynamicTables, + UpdateDynamicTableOp, + compare_dynamic_tables, +) + + +def _compare_dynamic_tables(autogen_context: AutogenContext, upgrade_ops, _): + dynamic_tables: DynamicTables | None = DynamicTables.extract(autogen_context.metadata) + if not dynamic_tables: + return + + assert autogen_context.connection + result = compare_dynamic_tables(autogen_context.connection, dynamic_tables) + upgrade_ops.ops.extend(result) + + +def render_dynamic_table(autogen_context: AutogenContext, op: DynamicTableOperation): + assert autogen_context.connection + dialect = autogen_context.connection.dialect + commands = op.to_sql(dialect) + return [f'op.execute("""{command}""")' for command in commands] + + +register_comparator_dispatcher(_compare_dynamic_tables, target="schema") +register_renderer_dispatcher( + CreateDynamicTableOp, UpdateDynamicTableOp, DropDynamicTableOp, fn=render_dynamic_table +) +register_rewriter_dispatcher(CreateDynamicTableOp, UpdateDynamicTableOp, DropDynamicTableOp) diff --git a/src/sqlalchemy_declarative_extensions/api.py b/src/sqlalchemy_declarative_extensions/api.py index 779ea52..9e96cda 100644 --- a/src/sqlalchemy_declarative_extensions/api.py +++ b/src/sqlalchemy_declarative_extensions/api.py @@ -6,6 +6,11 @@ from sqlalchemy.sql.schema import MetaData from sqlalchemy_declarative_extensions.database.base import Databases +from sqlalchemy_declarative_extensions.dialects.snowflake.dynamic_table import ( + DynamicTable, + DynamicTables, + dynamic_table_ddl, +) from sqlalchemy_declarative_extensions.function.base import Function, Functions from sqlalchemy_declarative_extensions.grant.base import Grants from sqlalchemy_declarative_extensions.procedure.base import Procedure, Procedures @@ -67,6 +72,7 @@ def declarative_database(base: T) -> T: raw_triggers = getattr(base, "triggers", None) raw_databases = getattr(base, "databases", None) raw_rows = getattr(base, "rows", None) + raw_dynamic_tables = getattr(base, "snowflake_dynamic_tables", None) metadata = getattr(base, "metadata", None) if metadata is None: # pragma: no cover @@ -83,6 +89,7 @@ def declarative_database(base: T) -> T: triggers=raw_triggers, databases=raw_databases, rows=raw_rows, + snowflake_dynamic_tables=raw_dynamic_tables, ) return base @@ -99,6 +106,7 @@ def declare_database( triggers: None | Iterable[Trigger] | Triggers = None, databases: None | Iterable[Database] | Databases = None, rows: None | Iterable[Row] | Rows = None, + snowflake_dynamic_tables: None | Iterable[DynamicTable] | DynamicTables = None, ): """Register declaratively specified database extension handlers. @@ -140,6 +148,7 @@ def declare_database( metadata.info["triggers"] = Triggers.coerce_from_unknown(triggers) metadata.info["databases"] = Databases.coerce_from_unknown(databases) metadata.info["rows"] = Rows.coerce_from_unknown(rows) + metadata.info["dynamic_tables"] = DynamicTables.coerce_from_unknown(snowflake_dynamic_tables) def register_sqlalchemy_events( @@ -154,6 +163,7 @@ def register_sqlalchemy_events( functions: bool | list[str] = False, triggers: bool | list[str] = False, rows: bool | list[str] = False, + snowflake_dynamic_tables: bool | list[str] = False, ): """Register handlers for supported object types into SQLAlchemy's event system. @@ -187,6 +197,7 @@ def register_sqlalchemy_events( functions=functions, triggers=triggers, rows=rows, + snowflake_dynamic_tables=snowflake_dynamic_tables, ) register_drop_events( @@ -198,6 +209,7 @@ def register_sqlalchemy_events( procedures=procedures, functions=functions, triggers=triggers, + snowflake_dynamic_tables=snowflake_dynamic_tables, ) @@ -213,6 +225,7 @@ def register_create_events( functions: bool | list[str] = False, triggers: bool | list[str] = False, rows: bool | list[str] = False, + snowflake_dynamic_tables: bool | list[str] = False, ): from sqlalchemy_declarative_extensions.database.ddl import database_ddl from sqlalchemy_declarative_extensions.function.ddl import function_ddl @@ -233,6 +246,7 @@ def register_create_events( concrete_triggers = metadata.info.get("triggers") or Triggers() concrete_databases = metadata.info.get("databases") or Databases() concrete_rows = metadata.info.get("rows") or Rows() + concrete_dynamic_tables = DynamicTables.extract(metadata) or DynamicTables() if databases: database_filter = databases if isinstance(databases, list) else None @@ -307,6 +321,14 @@ def register_create_events( rows_query(concrete_rows, row_filter), ) + if snowflake_dynamic_tables: + table_filter = snowflake_dynamic_tables if isinstance(snowflake_dynamic_tables, list) else None + event.listen( + metadata, + "after_create", + dynamic_table_ddl(concrete_dynamic_tables, table_filter), + ) + def register_drop_events( metadata: MetaData, @@ -318,6 +340,7 @@ def register_drop_events( procedures: bool | list[str] = False, functions: bool | list[str] = False, triggers: bool | list[str] = False, + snowflake_dynamic_tables: bool | list[str] = False, ): # Note grants and rows are (currently) omitted. Rows should handled by tables being dropped. # Grants should be handled by everything else being dropped. @@ -336,6 +359,7 @@ def register_drop_events( concrete_functions = metadata.info.get("functions") concrete_triggers = metadata.info.get("triggers") concrete_databases = metadata.info.get("databases") + concrete_dynamic_tables = metadata.info.get("dynamic_tables") if concrete_procedures and procedures: procedure_filter = procedures if isinstance(procedures, list) else None @@ -392,3 +416,11 @@ def register_drop_events( "after_drop", database_ddl(concrete_databases.are(), database_filter), ) + + if concrete_dynamic_tables and snowflake_dynamic_tables: + table_filter = snowflake_dynamic_tables if isinstance(snowflake_dynamic_tables, list) else None + event.listen( + metadata, + "before_drop", + dynamic_table_ddl(concrete_dynamic_tables.are(), table_filter), + ) diff --git a/src/sqlalchemy_declarative_extensions/dialects/snowflake/__init__.py b/src/sqlalchemy_declarative_extensions/dialects/snowflake/__init__.py index fbaec90..fd900bf 100644 --- a/src/sqlalchemy_declarative_extensions/dialects/snowflake/__init__.py +++ b/src/sqlalchemy_declarative_extensions/dialects/snowflake/__init__.py @@ -1,8 +1,20 @@ +from __future__ import annotations + +from sqlalchemy_declarative_extensions.dialects.snowflake.dynamic_table import ( + DynamicTable, + DynamicTables, + dynamic_table, + register_dynamic_table, +) from sqlalchemy_declarative_extensions.dialects.snowflake.role import Role from sqlalchemy_declarative_extensions.dialects.snowflake.schema import Schema from sqlalchemy_declarative_extensions.dialects.snowflake.view import View __all__ = [ + "DynamicTable", + "DynamicTables", + "dynamic_table", + "register_dynamic_table", "Role", "Schema", "View", diff --git a/src/sqlalchemy_declarative_extensions/dialects/snowflake/dynamic_table.py b/src/sqlalchemy_declarative_extensions/dialects/snowflake/dynamic_table.py new file mode 100644 index 0000000..45a10cd --- /dev/null +++ b/src/sqlalchemy_declarative_extensions/dialects/snowflake/dynamic_table.py @@ -0,0 +1,346 @@ +from __future__ import annotations + +import inspect +from dataclasses import dataclass, field, replace +from fnmatch import fnmatch +from typing import Any, Callable, Iterable, Sequence, TypeVar, Union + +from sqlalchemy import MetaData, text +from sqlalchemy.engine import Connection, Dialect +from sqlalchemy.sql import Select +from typing_extensions import Self + +from sqlalchemy_declarative_extensions.op import ExecuteOp +from sqlalchemy_declarative_extensions.sql import match_name, qualify_name +from sqlalchemy_declarative_extensions.sqlalchemy import HasMetaData, escape_params + +T = TypeVar("T") + + +def dynamic_table( + base, + *, + target_lag: str, + warehouse: str, +) -> Callable[[T], T]: + """Decorate a class to register a Snowflake Dynamic Table. + + Given an object with ``__tablename__``, optionally ``__table_args__``, + and ``__view__``, registers a DynamicTable. + + Arguments: + base: A declarative base object + target_lag: How stale the dynamic table's content can be (e.g. ``'1 minute'``) + warehouse: The warehouse used to refresh the dynamic table + """ + metadata = getattr(base, "metadata", None) + if metadata is None: + raise ValueError("Model must have a 'metadata' attribute.") + + def decorator(cls: T) -> T: + instance = DeclarativeDynamicTable(cls, target_lag=target_lag, warehouse=warehouse) + register_dynamic_table(base, instance) + return cls + + return decorator + + +def register_dynamic_table( + base_or_metadata: HasMetaData | MetaData, + dt: DynamicTable | DeclarativeDynamicTable, +): + """Register a dynamic table onto the given declarative base or MetaData.""" + if isinstance(base_or_metadata, MetaData): + metadata = base_or_metadata + else: + metadata = base_or_metadata.metadata + + if not metadata.info.get("dynamic_tables"): + metadata.info["dynamic_tables"] = DynamicTables() + metadata.info["dynamic_tables"].append(dt) + + +@dataclass +class DeclarativeDynamicTable: + cls: type + target_lag: str + warehouse: str + + @property + def name(self) -> str: + return self.cls.__tablename__ + + @property + def table_args(self): + return getattr(self.cls, "__table_args__", None) + + @property + def view_def(self) -> str | Select: + if inspect.isfunction(self.cls.__view__): + return self.cls.__view__() + return self.cls.__view__ + + @property + def schema(self) -> str | None: + table_args = self.table_args + if isinstance(table_args, dict): + return table_args.get("schema") + if isinstance(table_args, Iterable): + for table_arg in table_args: + if isinstance(table_arg, dict): + return table_arg.get("schema") + return None + + +@dataclass +class DynamicTable: + """Definition of a Snowflake Dynamic Table.""" + + name: str + definition: str | Select + target_lag: str + warehouse: str + schema: str | None = None + + @classmethod + def coerce_from_unknown(cls, unknown: Any) -> DynamicTable: + if isinstance(unknown, DynamicTable): + return cls( + name=unknown.name.upper(), + definition=unknown.definition, + target_lag=unknown.target_lag, + warehouse=unknown.warehouse.upper(), + schema=unknown.schema.upper() if unknown.schema else None, + ) + if isinstance(unknown, DeclarativeDynamicTable): + return cls( + name=unknown.name.upper(), + definition=unknown.view_def, + target_lag=unknown.target_lag, + warehouse=unknown.warehouse.upper(), + schema=unknown.schema.upper() if unknown.schema else None, + ) + raise NotImplementedError(f"Unsupported dynamic table source: {unknown}") + + @property + def qualified_name(self) -> str: + return qualify_name(self.schema, self.name) + + def compile_definition(self, dialect: Dialect | None = None) -> str: + if isinstance(self.definition, str): + return self.definition + return str( + self.definition.compile( + dialect=dialect, + compile_kwargs={"literal_binds": True}, + ) + ) + + def render_definition(self, conn: Connection) -> str: + compiled = self.compile_definition(conn.engine.dialect) + try: + import sqlglot + from sqlglot.optimizer.normalize import normalize + except ImportError: + raise ImportError("Dynamic table autogeneration requires the 'parse' extra.") + + return ( + escape_params( + normalize(sqlglot.parse_one(compiled, read="snowflake")).sql("snowflake") + ) + + ";" + ) + + def normalize(self, conn: Connection) -> Self: + definition = self.render_definition(conn) + return replace( + self, + name=self.name.upper(), + schema=self.schema.upper() if self.schema else None, + warehouse=self.warehouse.upper(), + definition=definition, + ) + + def to_sql_create(self, dialect: Dialect | None = None) -> list[str]: + definition = self.compile_definition(dialect).strip(";") + statement = ( + f"CREATE DYNAMIC TABLE {self.qualified_name}" + f" TARGET_LAG = '{self.target_lag}'" + f" WAREHOUSE = {self.warehouse}" + f" AS {definition};" + ) + return [statement] + + def to_sql_drop(self, dialect: Dialect | None = None) -> list[str]: + return [f"DROP DYNAMIC TABLE {self.qualified_name};"] + + def to_sql_update( + self, from_table: DynamicTable, dialect: Dialect | None = None + ) -> list[str]: + result = [] + result.extend(from_table.to_sql_drop(dialect)) + result.extend(self.to_sql_create(dialect)) + return result + + +@dataclass +class DynamicTables: + """Collection of Snowflake Dynamic Tables and associated comparison options.""" + + dynamic_tables: list[DynamicTable | DeclarativeDynamicTable] = field( + default_factory=list + ) + ignore_unspecified: bool = False + ignore: Iterable[str] = field(default_factory=set) + + @classmethod + def coerce_from_unknown( + cls, + unknown: None | Iterable[DynamicTable] | DynamicTables, + ) -> DynamicTables | None: + if isinstance(unknown, DynamicTables): + return unknown + if isinstance(unknown, Iterable): + return cls().are(*unknown) + return None + + @classmethod + def extract( + cls, + metadata: MetaData | list[MetaData] | list[MetaData | None] | None, + ) -> Self | None: + if not isinstance(metadata, Sequence): + metadata = [metadata] + + instances: list[Self] = [ + m.info["dynamic_tables"] + for m in metadata + if m and m.info.get("dynamic_tables") + ] + + if not instances: + return None + + tables: list[DynamicTable | DeclarativeDynamicTable] = [ + t for instance in instances for t in instance.dynamic_tables + ] + ignore: list[str] = [s for instance in instances for s in instance.ignore] + ignore_unspecified = instances[0].ignore_unspecified + + return cls( + dynamic_tables=tables, + ignore_unspecified=ignore_unspecified, + ignore=ignore, + ) + + def append(self, dynamic_table: DynamicTable | DeclarativeDynamicTable): + self.dynamic_tables.append(dynamic_table) + + def __iter__(self): + yield from self.dynamic_tables + + def are(self, *dynamic_tables: DynamicTable) -> Self: + return replace(self, dynamic_tables=list(dynamic_tables)) + + +@dataclass +class CreateDynamicTableOp(ExecuteOp): + dynamic_table: DynamicTable + + def reverse(self): + return DropDynamicTableOp(self.dynamic_table) + + def to_sql(self, dialect: Dialect | None = None) -> list[str]: + return self.dynamic_table.to_sql_create(dialect) + + +@dataclass +class UpdateDynamicTableOp(ExecuteOp): + from_dynamic_table: DynamicTable + dynamic_table: DynamicTable + + def reverse(self): + return UpdateDynamicTableOp( + from_dynamic_table=self.dynamic_table, + dynamic_table=self.from_dynamic_table, + ) + + def to_sql(self, dialect: Dialect | None = None) -> list[str]: + return self.dynamic_table.to_sql_update(self.from_dynamic_table, dialect) + + +@dataclass +class DropDynamicTableOp(ExecuteOp): + dynamic_table: DynamicTable + + def reverse(self): + return CreateDynamicTableOp(self.dynamic_table) + + def to_sql(self, dialect: Dialect | None = None) -> list[str]: + return self.dynamic_table.to_sql_drop(dialect) + + +DynamicTableOperation = Union[CreateDynamicTableOp, UpdateDynamicTableOp, DropDynamicTableOp] + + +def compare_dynamic_tables( + connection: Connection, + dynamic_tables: DynamicTables, +) -> list[DynamicTableOperation]: + from sqlalchemy_declarative_extensions.dialects.snowflake.query import ( + get_dynamic_tables_snowflake, + ) + + result: list[DynamicTableOperation] = [] + + concrete_defined: list[DynamicTable] = [ + DynamicTable.coerce_from_unknown(dt) for dt in dynamic_tables.dynamic_tables + ] + + by_name = {dt.qualified_name: dt for dt in concrete_defined} + expected_names = set(by_name) + + existing = get_dynamic_tables_snowflake(connection) + existing_by_name = {dt.qualified_name: dt for dt in existing} + existing_names = set(existing_by_name) + + new_names = expected_names - existing_names + removed_names = existing_names - expected_names + + for dt in concrete_defined: + normalized = dt.normalize(connection) + name = normalized.qualified_name + + if any(fnmatch(name, pattern) for pattern in dynamic_tables.ignore): + continue + + if name in new_names: + result.append(CreateDynamicTableOp(normalized)) + else: + existing_dt = existing_by_name[name] + normalized_existing = existing_dt.normalize(connection) + + if normalized_existing != normalized: + result.append(UpdateDynamicTableOp(normalized_existing, normalized)) + + if not dynamic_tables.ignore_unspecified: + for removed_name in removed_names: + if any(fnmatch(removed_name, pattern) for pattern in dynamic_tables.ignore): + continue + result.append(DropDynamicTableOp(existing_by_name[removed_name])) + + return result + + +def dynamic_table_ddl( + dynamic_tables: DynamicTables, table_filter: list[str] | None = None +): + def after_create(metadata: MetaData, connection: Connection, **_): + result = compare_dynamic_tables(connection, dynamic_tables) + for op in result: + if not match_name(op.dynamic_table.qualified_name, table_filter): + continue + for command in op.to_sql(connection.dialect): + connection.execute(text(command)) + + return after_create diff --git a/src/sqlalchemy_declarative_extensions/dialects/snowflake/query.py b/src/sqlalchemy_declarative_extensions/dialects/snowflake/query.py index 45562cf..e18d88d 100644 --- a/src/sqlalchemy_declarative_extensions/dialects/snowflake/query.py +++ b/src/sqlalchemy_declarative_extensions/dialects/snowflake/query.py @@ -69,6 +69,41 @@ def get_databases_snowflake(connection: Connection): } +def get_dynamic_tables_snowflake(connection: Connection): + from sqlalchemy_declarative_extensions.dialects.snowflake.dynamic_table import ( + DynamicTable, + ) + + query = text( + """ + SELECT table_schema, table_name, target_lag, warehouse, text + FROM information_schema.dynamic_tables + WHERE table_schema != 'INFORMATION_SCHEMA' + AND table_catalog = current_database() + """ + ) + + tables = [] + for row in connection.execute(query).fetchall(): + text_str: str = row.text + text_lower = text_str.lower() + warehouse_pos = text_lower.find("warehouse") + as_pos = text_lower.find(" as ", warehouse_pos) + definition = text_str[as_pos + 4:].strip() if as_pos != -1 else text_str + + schema = row.table_schema if row.table_schema != "PUBLIC" else None + tables.append( + DynamicTable( + name=row.table_name, + definition=definition, + target_lag=row.target_lag, + warehouse=row.warehouse, + schema=schema, + ) + ) + return tables + + def get_views_snowflake(connection: Connection): views_query = text( """