Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions src/sqlalchemy_declarative_extensions/alembic/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand Down Expand Up @@ -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
Expand Down
41 changes: 41 additions & 0 deletions src/sqlalchemy_declarative_extensions/alembic/dynamic_table.py
Original file line number Diff line number Diff line change
@@ -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)
32 changes: 32 additions & 0 deletions src/sqlalchemy_declarative_extensions/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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
Expand All @@ -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

Expand All @@ -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.

Expand Down Expand Up @@ -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(
Expand All @@ -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.

Expand Down Expand Up @@ -187,6 +197,7 @@ def register_sqlalchemy_events(
functions=functions,
triggers=triggers,
rows=rows,
snowflake_dynamic_tables=snowflake_dynamic_tables,
)

register_drop_events(
Expand All @@ -198,6 +209,7 @@ def register_sqlalchemy_events(
procedures=procedures,
functions=functions,
triggers=triggers,
snowflake_dynamic_tables=snowflake_dynamic_tables,
)


Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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,
Expand All @@ -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.
Expand All @@ -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
Expand Down Expand Up @@ -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),
)
Original file line number Diff line number Diff line change
@@ -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",
Expand Down
Loading
Loading