Skip to content
Merged
Show file tree
Hide file tree
Changes from 3 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
2 changes: 1 addition & 1 deletion README.md
Original file line number Diff line number Diff line change
Expand Up @@ -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 `<table>_temp` table, deduped there and merged into the destination, 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
Expand Down
87 changes: 51 additions & 36 deletions snowflake_utils/models/table.py
Original file line number Diff line number Diff line change
Expand Up @@ -253,39 +253,52 @@ 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:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Another potentially breaking change is this doesn't allow multiple COPY INTOs at the same time since they'd all reuse the same temp table name. I don't think it's relevant for our pipelines, but we could use a unique temp table name per run to fix it

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done in 358d827: unique temp table name per run, as suggested.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I just taught about a side effect of this - if the pod is killed while copying into the temp table, nobody cleans it up (and it's not tagged either, so we risk a leak). I think we should make the temp table TEMPORARY running all commands in the same connection?

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Follow-up in #36: the staging table is now TEMPORARY and the whole load runs in one session, so a killed pod no longer leaves a table in the schema (Snowflake discards it when that session expires). Thanks for the catch.

# 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,
)
Comment thread
cursor[bot] marked this conversation as resolved.
Comment thread
cursor[bot] marked this conversation as resolved.
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,
)
Comment thread
cursor[bot] marked this conversation as resolved.

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
Expand Down Expand Up @@ -331,16 +344,21 @@ def _merge(
primary_keys: list[str] = ["id"],
replication_keys: list[str] | None = None,
qualify: bool = False,
sync_tags: bool = True,
) -> 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"})
with connect() as connection:
connection.cursor().execute(f"drop table if exists {temp_table.fqn}")
copy_callable(temp_table, sync_tags=False)
if qualify:
with connect() as connection:
Expand All @@ -349,9 +367,6 @@ def _merge(

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)

Expand All @@ -364,7 +379,7 @@ def _merge(
temp_table, new_columns, old_columns, primary_keys
)
)
if self.table_structure:
if sync_tags and self.table_structure:
self.sync_tags(cursor)
temp_table.drop(cursor)

Expand Down
142 changes: 142 additions & 0 deletions tests/test_models.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
import inspect
import logging
import os
from datetime import datetime
Expand Down Expand Up @@ -407,6 +408,147 @@ 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}


@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 "drop table if exists PUBLIC.PYTEST_temp" 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_sync_tags.assert_not_called()


@patch("snowflake_utils.settings.connect")
def test_single_column_update(mock_connect):
mock_cursor = make_mock_cursor()
Expand Down
Loading