diff --git a/README.md b/README.md index 0266f5f..5f0e392 100644 --- a/README.md +++ b/README.md @@ -68,7 +68,7 @@ Available (public) methods: - whether the table should be fully refreshed - a list of target columns, if not all the columns are present in the file to be loaded (or not all need to be written) - whether to sync tags (provided in the table structure) to the table/columns - - whether to dedupe the loaded data (qualify). If this is used, a list of primary keys (and optionally of replication keys) should also be provided. If only the primary keys are supplied, those will also be used to determine which records are kept (non deterministic). Unless `full_refresh` is set, the data is copied to a uniquely named `_temp_` table, deduped there and merged into the destination (updating only the loaded columns, and only rows that are not older by the replication keys), so the destination never holds duplicates visible to readers. With `full_refresh` the table is copied into directly and then rebuilt with the qualify. + - whether to dedupe the loaded data (qualify). If this is used, a list of primary keys (and optionally of replication keys) should also be provided. If only the primary keys are supplied, those will also be used to determine which records are kept (non deterministic). Unless `full_refresh` is set, the data is copied, in a single session, to a `TEMPORARY` table (`
_temp_`, visible only to that session and discarded with it), deduped there and merged into the destination (updating only the loaded columns, and only rows that are not older by the replication keys), so the destination never holds duplicates visible to readers. With `full_refresh` the table is copied into directly and then rebuilt with the qualify. - a stage parameter to use an existing stage instead of creating a temporary one - *create_table*: runs the create table statement, with optional full refresh to recreate an existing table. - *setup_file_format*: given a file format object, creates the corresponding resource in Snowflake diff --git a/snowflake_utils/models/table.py b/snowflake_utils/models/table.py index 582f420..f644b96 100644 --- a/snowflake_utils/models/table.py +++ b/snowflake_utils/models/table.py @@ -1,10 +1,14 @@ import logging from collections import defaultdict +from collections.abc import Iterator +from contextlib import contextmanager +from contextvars import ContextVar from functools import partial from typing import ClassVar from uuid import uuid4 from pydantic import BaseModel, Field +from snowflake.connector import SnowflakeConnection from snowflake.connector.cursor import SnowflakeCursor from ..queries import execute_statement @@ -21,6 +25,18 @@ from .file_format import FileFormat, InlineFileFormat from .table_structure import TableStructure +_session: ContextVar[SnowflakeConnection | None] = ContextVar("_session", default=None) + + +@contextmanager +def _connection() -> Iterator[SnowflakeConnection]: + """Reuse the session opened by `Table._merge`, otherwise open a new connection.""" + if (shared := _session.get()) is not None: + yield shared + else: + with connect() as connection: + yield connection + class Table(BaseModel): name: str @@ -30,6 +46,7 @@ class Table(BaseModel): database: str | None = None include_metadata: list[MetadataColumn] = Field(default_factory=list) enable_schema_evolution: bool = False + temporary: bool = False existing_column_tags: dict[str, dict[str, str]] | None = None existing_table_tags: dict[str, str] | None = None _file_format: FileFormat | None = None @@ -122,8 +139,14 @@ def get_create_table_statement( ) -> str: logging.debug(f"Creating table: {self.fqn}") copy_grants_clause = " COPY GRANTS" if copy_grants and full_refresh else "" + kind = "TEMPORARY TABLE" if self.temporary else "TABLE" + create = ( + f"CREATE OR REPLACE {kind}" + if full_refresh + else f"CREATE {kind} IF NOT EXISTS" + ) if self.table_structure: - return f"{'CREATE OR REPLACE TABLE' if full_refresh else 'CREATE TABLE IF NOT EXISTS'} {self.fqn}{copy_grants_clause} ({self.table_structure.parsed_columns})" + return f"{create} {self.fqn}{copy_grants_clause} ({self.table_structure.parsed_columns})" else: template = """ARRAY_AGG( OBJECT_CONSTRUCT( @@ -150,7 +173,7 @@ def get_create_table_statement( stage_query = f"LOCATION => '@{self.stage}'" return f""" - {"CREATE OR REPLACE TABLE" if full_refresh else "CREATE TABLE IF NOT EXISTS"} {self.fqn}{copy_grants_clause} + {create} {self.fqn}{copy_grants_clause} USING TEMPLATE ( SELECT {template} FROM TABLE( @@ -168,7 +191,7 @@ def bulk_insert( records, full_refresh: bool = False, ) -> None: - with connect() as connection: + with _connection() as connection: cursor = connection.cursor() _execute_statement = partial(execute_statement, cursor) _execute_statement(self.get_create_schema_statement()) @@ -198,7 +221,7 @@ def _copy( create_table: bool = True, copy_grants: bool = True, ) -> None: - with connect() as connection: + with _connection() as connection: cursor = connection.cursor() execute = self.setup_connection( path, storage_integration, cursor, file_format, stage @@ -299,7 +322,7 @@ def copy_callable(table: Table, sync_tags: bool) -> None: ) if not qualify: return result - with connect() as connection: + with _connection() as connection: cursor = connection.cursor() self.qualify( cursor=cursor, @@ -356,28 +379,30 @@ def _merge( sync_tags: bool = True, target_columns: list[str] | None = None, ) -> None: + # one session for the whole load: the temp table only exists inside it and + # Snowflake drops it when the session ends, even if the process is killed with connect() as connection: - cursor = connection.cursor() - if not self.exists(cursor): - copy_callable(self, sync_tags=sync_tags) + token = _session.set(connection) + try: + cursor = connection.cursor() + if not self.exists(cursor): + copy_callable(self, sync_tags=sync_tags) + if qualify: + self.qualify(cursor, primary_keys, replication_keys) + if sync_tags and self.table_structure: + self.sync_tags(cursor) + return None + + temp_table = self.model_copy( + update={ + "name": f"{self.name}_temp_{uuid4().hex[:8]}", + "temporary": True, + } + ) + copy_callable(temp_table, sync_tags=False) if qualify: - self.qualify(cursor, primary_keys, replication_keys) - if sync_tags and self.table_structure: - self.sync_tags(cursor) - return None - - temp_table = self.model_copy( - update={"name": f"{self.name}_temp_{uuid4().hex[:8]}"} - ) - try: - copy_callable(temp_table, sync_tags=False) - if qualify: - with connect() as connection: - cursor = connection.cursor() temp_table.qualify(cursor, primary_keys, replication_keys) - with connect() as connection: - cursor = connection.cursor() old_columns = {x.name: x.data_type for x in self.get_columns(cursor)} new_columns = temp_table.get_columns(cursor) @@ -410,10 +435,8 @@ def _merge( if sync_tags and self.table_structure: self.sync_tags(cursor) temp_table.drop(cursor) - except BaseException: - with connect() as connection: - connection.cursor().execute(f"drop table if exists {temp_table.fqn}") - raise + finally: + _session.reset(token) def merge( self, @@ -501,7 +524,7 @@ def qualify( ) return cursor.execute( f""" - create or replace table {self.fqn} as ( + create or replace {"temporary " if self.temporary else ""}table {self.fqn} as ( select * from {self.fqn} qualify row_number() over (partition by {qualify_partition} order by {qualify_order}) = 1 ) diff --git a/tests/test_models.py b/tests/test_models.py index d0576df..38b70ba 100644 --- a/tests/test_models.py +++ b/tests/test_models.py @@ -17,6 +17,7 @@ TableStructure, ) from snowflake_utils.models.column import MetadataColumn +from snowflake_utils.models.table import _connection, _session test_table_schema = TableStructure( columns={ @@ -538,7 +539,7 @@ def test_copy_into_qualify_never_rebuilds_live_table( for s in statements ) assert any( - s.lower().startswith("create or replace table public.pytest_temp") + s.lower().startswith("create or replace temporary table public.pytest_temp_") for s in statements ) assert any( @@ -628,15 +629,47 @@ def test_copy_into_qualify_uses_a_new_temp_table_per_run( assert len(temp_names) == 2 +@patch.object(Table, "drop") +@patch.object(Table, "get_columns") +@patch.object(Table, "exists", return_value=True) +@patch.object(Table, "_copy", autospec=True) +def test_copy_into_qualify_runs_everything_in_one_session( + mock_copy, mock_exists, mock_get_columns, mock_drop +): + mock_get_columns.return_value = [Column(name="id", data_type="text")] + sessions = [] + + def copy_in_session(self, *args, **kwargs): + with _connection() as connection: + sessions.append((self.temporary, connection)) + + mock_copy.side_effect = copy_in_session + mock_conn = make_mock_conn() + with patch( + "snowflake_utils.models.table.connect", return_value=mock_conn + ) as connect: + test_table.copy_into( + path=path, + file_format=parquet_file_format, + storage_integration=storage_integration, + primary_keys=["id"], + qualify=True, + ) + + connect.assert_called_once() + assert sessions == [(True, mock_conn)] + mock_conn.__exit__.assert_called_once() + + @patch.object(Table, "get_columns") @patch.object(Table, "exists", return_value=True) @patch.object(Table, "_copy", side_effect=RuntimeError("copy failed")) -def test_copy_into_qualify_drops_temp_table_when_the_load_fails( +def test_copy_into_qualify_failure_closes_the_session_without_merging( mock_copy, mock_exists, mock_get_columns ): mock_cursor = make_mock_cursor() - with patch("snowflake_utils.models.table.connect") as mock_connect: - mock_connect.return_value = make_mock_conn(cursor=mock_cursor) + mock_conn = make_mock_conn(cursor=mock_cursor) + with patch("snowflake_utils.models.table.connect", return_value=mock_conn): with pytest.raises(RuntimeError, match="copy failed"): test_table.copy_into( path=path, @@ -646,12 +679,41 @@ def test_copy_into_qualify_drops_temp_table_when_the_load_fails( qualify=True, ) - statements = executed_statements(mock_cursor) - assert any( - re.fullmatch(r"drop table if exists PUBLIC\.PYTEST_temp_[0-9a-f]{8}", s) - for s in statements + mock_conn.__exit__.assert_called_once() + assert _session.get() is None + assert not any( + s.lower().startswith("merge into") for s in executed_statements(mock_cursor) + ) + + +@pytest.mark.parametrize( + "full_refresh, expected", + [ + (False, "CREATE TEMPORARY TABLE IF NOT EXISTS PUBLIC.PYTEST_temp_1 ("), + (True, "CREATE OR REPLACE TEMPORARY TABLE PUBLIC.PYTEST_temp_1 ("), + ], +) +def test_temporary_table_create_statement(full_refresh, expected): + temp = test_table.model_copy(update={"name": "PYTEST_temp_1", "temporary": True}) + + statement = temp.get_create_table_statement(full_refresh, copy_grants=False) + + assert statement.startswith(expected) + assert "TEMPORARY" not in test_table.get_create_table_statement() + + +@pytest.mark.parametrize("temporary", [True, False]) +def test_qualify_recreates_the_table_with_the_same_kind(temporary): + mock_cursor = make_mock_cursor() + table = test_table.model_copy(update={"temporary": temporary}) + + table.qualify(mock_cursor, ["id"], None) + + statement = " ".join(mock_cursor.execute.call_args.args[0].split()).lower() + expected = ( + "create or replace temporary table" if temporary else "create or replace table" ) - assert not any(s.lower().startswith("merge into") for s in statements) + assert statement.startswith(expected) @patch("snowflake_utils.settings.connect")