diff --git a/README.md b/README.md
index cb5331d..0266f5f 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 perform a qualify on the table after loading the data. 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).
+ - 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.
- 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/column.py b/snowflake_utils/models/column.py
index fa9a19f..0b1002e 100644
--- a/snowflake_utils/models/column.py
+++ b/snowflake_utils/models/column.py
@@ -25,6 +25,18 @@ def tmp(x: str) -> str:
)
+def _is_newer_or_equal(keys: list[str]) -> str:
+ *head, last = (k.upper() for k in keys)
+ expr = f'(dest."{last}" is null or tmp."{last}" >= dest."{last}")'
+ for key in reversed(head):
+ dest, tmp = f'dest."{key}"', f'tmp."{key}"'
+ expr = (
+ f"(({dest} is null and {tmp} is not null) or {tmp} > {dest}"
+ f" or (equal_null({tmp}, {dest}) and {expr}))"
+ )
+ return expr
+
+
def _inserts(columns: list[Column], old_columns: dict[str, str]) -> str:
return ",".join(
_possibly_cast(f'tmp."{c.name}"', old_columns.get(c.name), c.data_type)
diff --git a/snowflake_utils/models/table.py b/snowflake_utils/models/table.py
index 2644019..582f420 100644
--- a/snowflake_utils/models/table.py
+++ b/snowflake_utils/models/table.py
@@ -2,13 +2,21 @@
from collections import defaultdict
from functools import partial
from typing import ClassVar
+from uuid import uuid4
from pydantic import BaseModel, Field
from snowflake.connector.cursor import SnowflakeCursor
from ..queries import execute_statement
from ..settings import SnowflakeSettings, connect, governance_settings
-from .column import Column, MetadataColumn, _inserts, _matched, _type_cast
+from .column import (
+ Column,
+ MetadataColumn,
+ _inserts,
+ _is_newer_or_equal,
+ _matched,
+ _type_cast,
+)
from .enums import MatchByColumnName, TagLevel
from .file_format import FileFormat, InlineFileFormat
from .table_structure import TableStructure
@@ -253,40 +261,54 @@ def copy_into(
{files_clause}
{self._include_metadata()}
"""
- if qualify:
- self._copy(
- copy_query,
- path,
- file_format,
- storage_integration,
- full_refresh,
- sync_tags,
- stage,
- create_table,
- copy_grants,
- )
- with connect() as connection:
- cursor = connection.cursor()
- self.qualify(
- cursor=cursor,
- primary_keys=primary_keys,
- replication_keys=replication_keys,
+ if qualify and not full_refresh:
+ # dedupe in a temp table so the live table never holds duplicates
+ def copy_callable(table: Table, sync_tags: bool) -> None:
+ return table.copy_into(
+ path=path,
+ file_format=file_format,
+ storage_integration=storage_integration,
+ match_by_column_name=match_by_column_name,
+ target_columns=target_columns,
+ sync_tags=sync_tags,
+ stage=stage,
+ files=files,
+ create_table=create_table or table is not self,
+ copy_grants=copy_grants,
)
- if sync_tags and self.table_structure:
- self.sync_tags(cursor)
- else:
- return self._copy(
- copy_query,
- path,
- file_format,
- storage_integration,
- full_refresh,
- sync_tags,
- stage,
- create_table,
- copy_grants,
+
+ return self._merge(
+ copy_callable,
+ primary_keys,
+ replication_keys,
+ qualify=True,
+ sync_tags=sync_tags,
+ target_columns=target_columns,
)
+ result = self._copy(
+ copy_query,
+ path,
+ file_format,
+ storage_integration,
+ full_refresh,
+ sync_tags,
+ stage,
+ create_table,
+ copy_grants,
+ )
+ if not qualify:
+ return result
+ with connect() as connection:
+ cursor = connection.cursor()
+ self.qualify(
+ cursor=cursor,
+ primary_keys=primary_keys,
+ replication_keys=replication_keys,
+ )
+ if sync_tags and self.table_structure:
+ self.sync_tags(cursor)
+
def create_table(
self, full_refresh: bool, execute_statement: callable, copy_grants: bool = True
) -> None:
@@ -331,42 +353,67 @@ def _merge(
primary_keys: list[str] = ["id"],
replication_keys: list[str] | None = None,
qualify: bool = False,
+ sync_tags: bool = True,
+ target_columns: list[str] | None = None,
) -> None:
with connect() as connection:
cursor = connection.cursor()
if not self.exists(cursor):
- copy_callable(self, sync_tags=True)
+ 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"})
- copy_callable(temp_table, sync_tags=False)
- if qualify:
+ 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()
- temp_table.qualify(cursor, primary_keys, replication_keys)
-
- with connect() as connection:
- cursor = connection.cursor()
- cursor.execute(
- self.get_create_table_statement(full_refresh=False, copy_grants=True)
- )
- old_columns = {x.name: x.data_type for x in self.get_columns(cursor)}
- new_columns = temp_table.get_columns(cursor)
-
- for column in new_columns:
- if column.name not in old_columns:
- self.add_column(cursor, column)
-
- cursor.execute(
- self._merge_statement(
- temp_table, new_columns, old_columns, primary_keys
+ old_columns = {x.name: x.data_type for x in self.get_columns(cursor)}
+ new_columns = temp_table.get_columns(cursor)
+
+ for column in new_columns:
+ if column.name not in old_columns:
+ self.add_column(cursor, column)
+
+ loaded = {
+ c.upper()
+ for c in [
+ *(target_columns or []),
+ *primary_keys,
+ *(m.name for m in self.include_metadata),
+ ]
+ }
+ merge_columns = [
+ c
+ for c in new_columns
+ if not target_columns or c.name.upper() in loaded
+ ]
+ cursor.execute(
+ self._merge_statement(
+ temp_table,
+ merge_columns,
+ old_columns,
+ primary_keys,
+ replication_keys if qualify else None,
+ )
)
- )
- if self.table_structure:
- self.sync_tags(cursor)
- temp_table.drop(cursor)
+ 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
def merge(
self,
@@ -467,10 +514,14 @@ def _merge_statement(
columns: list[Column],
old_columns: dict[str, str],
primary_keys: list[str],
+ replication_keys: list[str] | None = None,
) -> str:
pkes = " and ".join(
f'dest."{c.upper()}" = tmp."{c.upper()}"' for c in primary_keys
)
+ update_condition = (
+ f" and {_is_newer_or_equal(replication_keys)}" if replication_keys else ""
+ )
matched = _matched(columns, old_columns)
column_names = ",".join(f'"{c.name}"' for c in columns)
inserts = _inserts(columns, old_columns)
@@ -483,7 +534,7 @@ def _merge_statement(
merge into {self.fqn} as dest
using {temp_table.fqn} tmp
ON {pkes}
- when matched then update set {matched}
+ when matched{update_condition} then update set {matched}
when not matched then insert ({column_names}) VALUES ({inserts})
"""
diff --git a/tests/test_models.py b/tests/test_models.py
index 41abe09..d0576df 100644
--- a/tests/test_models.py
+++ b/tests/test_models.py
@@ -1,5 +1,7 @@
+import inspect
import logging
import os
+import re
from datetime import datetime
from unittest.mock import MagicMock, patch
@@ -407,6 +409,251 @@ def test_merge(mock_merge, mock_copy):
}
+@patch.object(Table, "_copy")
+@patch.object(Table, "_merge")
+def test_copy_into_qualify_merges_instead_of_copying_into_live_table(
+ mock_merge, mock_copy
+):
+ test_table.copy_into(
+ path=path,
+ file_format=parquet_file_format,
+ storage_integration=storage_integration,
+ primary_keys=["id"],
+ qualify=True,
+ sync_tags=True,
+ )
+
+ mock_copy.assert_not_called()
+ mock_merge.assert_called_once()
+ _, kwargs = mock_merge.call_args
+ assert kwargs == {"qualify": True, "sync_tags": True, "target_columns": None}
+
+
+@patch.object(Table, "qualify")
+@patch.object(Table, "_copy")
+@patch.object(Table, "_merge")
+def test_copy_into_qualify_full_refresh_keeps_copy_then_qualify(
+ mock_merge, mock_copy, mock_qualify
+):
+ with patch("snowflake_utils.models.table.connect") as mock_connect:
+ mock_connect.return_value = make_mock_conn()
+ test_table.copy_into(
+ path=path,
+ file_format=parquet_file_format,
+ storage_integration=storage_integration,
+ primary_keys=["id"],
+ qualify=True,
+ full_refresh=True,
+ )
+
+ mock_merge.assert_not_called()
+ mock_copy.assert_called_once()
+ mock_qualify.assert_called_once()
+
+
+def copy_args(call) -> dict:
+ return inspect.signature(Table._copy).bind(*call.args, **call.kwargs).arguments
+
+
+@pytest.mark.parametrize("create_table", [True, False])
+@patch.object(Table, "_copy", autospec=True)
+@patch.object(Table, "_merge")
+def test_copy_into_qualify_always_creates_temp_table(
+ mock_merge, mock_copy, create_table
+):
+ test_table.copy_into(
+ path=path,
+ file_format=parquet_file_format,
+ storage_integration=storage_integration,
+ target_columns=["id"],
+ qualify=True,
+ create_table=create_table,
+ )
+ copy_callable = mock_merge.call_args.args[0]
+
+ copy_callable(
+ test_table.model_copy(update={"name": "PYTEST_temp"}), sync_tags=False
+ )
+ copy_callable(test_table, sync_tags=True)
+
+ temp_call, live_call = (copy_args(c) for c in mock_copy.call_args_list)
+ assert temp_call["sync_tags"] is False and temp_call["create_table"] is True
+ assert live_call["sync_tags"] is True and live_call["create_table"] is create_table
+ assert "COPY INTO PUBLIC.PYTEST_temp (id)" in temp_call["query"].replace("\n", " ")
+
+
+@patch.object(Table, "drop")
+@patch.object(Table, "get_columns")
+@patch.object(Table, "exists", return_value=True)
+@patch.object(Table, "_copy")
+def test_copy_into_qualify_existing_table_without_structure(
+ mock_copy, mock_exists, mock_get_columns, mock_drop
+):
+ mock_get_columns.return_value = [Column(name="id", data_type="integer")]
+ mock_cursor = make_mock_cursor()
+ inferred_table = Table(name="PYTEST_INFERRED", schema_name="PUBLIC")
+ with patch("snowflake_utils.models.table.connect") as mock_connect:
+ mock_connect.return_value = make_mock_conn(cursor=mock_cursor)
+ inferred_table.copy_into(
+ path=path,
+ file_format=parquet_file_format,
+ storage_integration=storage_integration,
+ primary_keys=["id"],
+ qualify=True,
+ )
+
+ statements = [
+ " ".join(c.args[0].split()) for c in mock_cursor.execute.call_args_list
+ ]
+ assert any(
+ s.lower().startswith("merge into public.pytest_inferred") for s in statements
+ )
+
+
+@patch.object(Table, "sync_tags")
+@patch.object(Table, "drop")
+@patch.object(Table, "get_columns")
+@patch.object(Table, "exists", return_value=True)
+@patch.object(Table, "_copy")
+def test_copy_into_qualify_never_rebuilds_live_table(
+ mock_copy, mock_exists, mock_get_columns, mock_drop, mock_sync_tags
+):
+ mock_get_columns.return_value = [Column(name="id", data_type="integer")]
+ 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)
+ test_table.copy_into(
+ path=path,
+ file_format=parquet_file_format,
+ storage_integration=storage_integration,
+ primary_keys=["id"],
+ qualify=True,
+ )
+
+ statements = [
+ " ".join(c.args[0].split()) for c in mock_cursor.execute.call_args_list
+ ]
+ assert not any(
+ s.lower().startswith("create or replace table public.pytest ")
+ for s in statements
+ )
+ assert any(
+ s.lower().startswith("create or replace table public.pytest_temp")
+ for s in statements
+ )
+ assert any(
+ s.lower().startswith("merge into public.pytest as dest") for s in statements
+ )
+ # the live table is only ever written by the MERGE: COPY targets the temp table
+ assert "PUBLIC.PYTEST_temp_" in mock_copy.call_args.args[0]
+ mock_drop.assert_called_once()
+ mock_sync_tags.assert_not_called()
+
+
+def run_qualified_copy_into(table: Table, **kwargs) -> MagicMock:
+ 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)
+ table.copy_into(
+ path=path,
+ file_format=parquet_file_format,
+ storage_integration=storage_integration,
+ primary_keys=["id"],
+ qualify=True,
+ **kwargs,
+ )
+ return mock_cursor
+
+
+def executed_statements(mock_cursor: MagicMock) -> list[str]:
+ return [" ".join(c.args[0].split()) for c in mock_cursor.execute.call_args_list]
+
+
+@patch.object(Table, "drop")
+@patch.object(Table, "get_columns")
+@patch.object(Table, "exists", return_value=True)
+@patch.object(Table, "_copy")
+def test_copy_into_qualify_only_updates_loaded_columns(
+ mock_copy, mock_exists, mock_get_columns, mock_drop
+):
+ mock_get_columns.return_value = [
+ Column(name=n, data_type="text") for n in ("id", "name", "last_name")
+ ]
+
+ statements = executed_statements(
+ run_qualified_copy_into(test_table, target_columns=["name"])
+ )
+
+ merge = next(s for s in statements if s.lower().startswith("merge into"))
+ update_clause = merge.split("when matched")[1].split("when not matched")[0]
+ assert 'dest."name"' in update_clause and 'dest."id"' in update_clause
+ assert "last_name" not in update_clause
+
+
+def test_merge_statement_updates_only_newer_rows_with_replication_keys():
+ columns = [Column(name="id", data_type="text")]
+ temp = test_table.model_copy(update={"name": "PYTEST_temp"})
+
+ with_keys = " ".join(
+ test_table._merge_statement(
+ temp, columns, {}, ["id"], ["updated_at", "etl_file_ingested_at"]
+ ).split()
+ )
+ without_keys = " ".join(
+ test_table._merge_statement(temp, columns, {}, ["id"]).split()
+ )
+
+ condition = with_keys.split("when matched")[1].split("then update")[0]
+ assert 'tmp."UPDATED_AT" > dest."UPDATED_AT"' in condition
+ assert 'tmp."ETL_FILE_INGESTED_AT" >= dest."ETL_FILE_INGESTED_AT"' in condition
+ assert "when matched then update" in without_keys
+
+
+@patch.object(Table, "drop")
+@patch.object(Table, "get_columns")
+@patch.object(Table, "exists", return_value=True)
+@patch.object(Table, "_copy")
+def test_copy_into_qualify_uses_a_new_temp_table_per_run(
+ mock_copy, mock_exists, mock_get_columns, mock_drop
+):
+ mock_get_columns.return_value = [Column(name="id", data_type="text")]
+
+ run_qualified_copy_into(test_table)
+ run_qualified_copy_into(test_table)
+
+ temp_names = {
+ re.search(r"PYTEST_temp_[0-9a-f]{8}", c.args[0]).group()
+ for c in mock_copy.call_args_list
+ }
+ assert len(temp_names) == 2
+
+
+@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(
+ 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)
+ with pytest.raises(RuntimeError, match="copy failed"):
+ test_table.copy_into(
+ path=path,
+ file_format=parquet_file_format,
+ storage_integration=storage_integration,
+ primary_keys=["id"],
+ 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
+ )
+ assert not any(s.lower().startswith("merge into") for s in statements)
+
+
@patch("snowflake_utils.settings.connect")
def test_single_column_update(mock_connect):
mock_cursor = make_mock_cursor()