Skip to content
Merged
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
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 uniquely named `<table>_temp_<id>` 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
Expand Down
12 changes: 12 additions & 0 deletions snowflake_utils/models/column.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
167 changes: 109 additions & 58 deletions snowflake_utils/models/table.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:

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,
target_columns=target_columns,
)
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
) -> None:
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand All @@ -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})
"""

Expand Down
Loading
Loading