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 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.
- 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 (`<table>_temp_<id>`, 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
Expand Down
79 changes: 51 additions & 28 deletions snowflake_utils/models/table.py
Original file line number Diff line number Diff line change
@@ -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
Expand All @@ -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
Expand All @@ -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
Expand Down Expand Up @@ -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(
Expand All @@ -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(
Expand All @@ -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())
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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)

Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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
)
Expand Down
80 changes: 71 additions & 9 deletions tests/test_models.py
Original file line number Diff line number Diff line change
Expand Up @@ -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={
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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,
Expand All @@ -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")
Expand Down
Loading