Skip to content

Instantly share code, notes, and snippets.

@philerooski
Created June 15, 2026 13:14
Show Gist options
  • Select an option

  • Save philerooski/63c96ec4015b57319b3c5d8e20c68e1c to your computer and use it in GitHub Desktop.

Select an option

Save philerooski/63c96ec4015b57319b3c5d8e20c68e1c to your computer and use it in GitHub Desktop.
Load RDS snapshot data from S3 into Snowflake with column-level masking for raw_table_read_censored
"""
Load RDS snapshot data from S3 via Snowflake external stage into tables.
This script supports two modes:
1. Bootstrap mode (--bootstrap-stack): Creates a new schema, external stage,
file format, and grants privileges before loading data.
2. Manual mode: Loads data into an existing schema with pre-configured stage.
The script dynamically discovers all data types from the S3 stage URL,
creates tables using INFER_SCHEMA from Parquet files, loads the data,
and manages database role privileges. Sensitive columns are masked via
Snowflake masking policies so that RAW_TABLE_READ_CENSORED can read all
tables but sees NULL for security-sensitive fields; RAW_TABLE_READ retains
full visibility. All operations are logged to LOAD_LOG.
See `python load_snapshot_data.py --help`
"""
import snowflake.connector
from typing import Optional
import sys
# Global configuration constants
DEFAULT_DATABASE = "SYNAPSE_RDS_SNAPSHOT"
STORAGE_INTEGRATION = "synapse_snapshot_poc"
FILE_FORMAT_NAME = "parquet_ff"
DB_ROLE_RAW_TABLE_READ = "raw_table_read"
DB_ROLE_RAW_TABLE_READ_CENSORED = "raw_table_read_censored"
# Tables that have a schema inferred as BINARY, but ought to be VARIANT
TABLES_WITH_VARIANT_COLUMNS = {
"ASYNCH_JOB_STATUS",
"CHALLENGE_TEAM",
"DATA_ACCESS_SUBMISSION_STATUS",
"DISCUSSION_THREAD",
"DOWNLOAD_ORDER",
"EVALUATION",
"EVALUATION_SUBMISSION",
"MESSAGE_TO_USER",
"MULTIPART_UPLOAD",
"MULTIPART_UPLOAD_PART_STATE",
"OAUTH_AUTHORIZATION_CODE",
"QUIZ_RESPONSE",
"RESEARCH_PROJECT",
"STATISTICS_MONTHLY_STATUS",
"SUBSTATUS_ANNOTATIONS_BLOB",
"TABLE_STATUS",
"USER_GROUP",
"V2_WIKI_MARKDOWN",
"V2_WIKI_OWNERS",
"VERIFICATION_STATE",
}
# Tables that have actual binary/compressed data and should keep BINARY columns
TABLES_WITH_TRUE_BINARY_DATA = {
"ACCESS_REQUIREMENT_REVISION",
"ACTIVITY",
"CHALLENGE",
"DATA_ACCESS_REQUEST",
"DATA_ACCESS_SUBMISSION",
"MEMBERSHIP_INVITATION_SUBMISSION",
"MEMBERSHIP_REQUEST_SUBMISSION",
"NODE_REVISION",
"PERSONAL_ACCESS_TOKEN",
"TEAM",
"USER_PROFILE",
"VERIFICATION_SUBMISSION",
}
# Columns that must be masked for RAW_TABLE_READ_CENSORED.
# Maps table name (upper) -> {column name (upper) -> masking policy name}.
# Column types and policy names were verified against PROD_588.
MASKED_COLUMNS: dict[str, dict[str, str]] = {
"CREDENTIAL": {
"PASS_HASH": "pii_mask_text",
},
"OAUTH_ACCESS_TOKEN": {
"TOKEN_ID": "pii_mask_text",
},
"OAUTH_AUTHORIZATION_CODE": {
"AUTHORIZATION_REQUEST": "pii_mask_binary",
"AUTH_CODE": "pii_mask_text",
},
"OAUTH_CLIENT": {
"SECRET_HASH": "pii_mask_text",
},
"OAUTH_REFRESH_TOKEN": {
"TOKEN_HASH": "pii_mask_text",
},
"OAUTH_SECTOR_IDENTIFIER": {
"SECRET": "pii_mask_text",
},
"OTP_RECOVERY_CODE": {
"CODE_HASH": "pii_mask_text",
},
"OTP_SECRET": {
"SECRET": "pii_mask_text",
},
"PERSONAL_ACCESS_TOKEN": {
"SCOPES": "pii_mask_binary",
"CLAIMS": "pii_mask_binary",
},
}
def derive_data_type(prefix: str, prefix_base: str) -> str:
"""
Derive data type (table name) from prefix.
Example: 'dev566/dev566.NODE/1/' -> 'NODE'
Args:
prefix: The prefix string to parse (e.g. from `s3_prefixes.txt`)
prefix_base: The base prefix used in S3 keys/stage paths (e.g. 'dev566')
Returns:
The data type/table name extracted from the prefix
"""
# Split by '<prefix_base>.' and take the second part
parts = prefix.split(f"{prefix_base}.", 1)
if len(parts) < 2:
raise ValueError(f"Invalid prefix format: {prefix}")
# Split by '/' and take the first part
data_type = parts[1].split("/")[0]
return data_type
def derive_stage_path(data_type: str, prefix_base: str) -> str:
"""
Derive stage path from data type.
Args:
data_type: The data type/table name
prefix_base: The base prefix used in S3 keys/stage paths
Returns:
Stage path string formed as '<prefix_base>.<DATA_TYPE>/1/'
"""
return f"{prefix_base}.{data_type}/1/"
def log_operation(
cursor,
prefix: str,
data_type: str,
stage_path: str,
phase: str,
status: str,
sql_text: Optional[str] = None,
error_msg: Optional[str] = None,
):
"""
Log an operation to the LOAD_LOG table.
Args:
cursor: Snowflake cursor for executing queries
prefix: The prefix being processed
data_type: The derived data type
stage_path: The stage path
phase: The current phase (START, CREATE_TABLE, COPY_DATA, ERROR)
status: Status of the operation (OK, RUN, FAIL)
sql_text: Optional SQL text being executed
error_msg: Optional error message
"""
log_sql = """
INSERT INTO LOAD_LOG (PREFIX, DATA_TYPE, STAGE_PATH, PHASE, STATUS, SQL_TEXT, ERROR_MESSAGE)
VALUES (%s, %s, %s, %s, %s, %s, %s)
"""
cursor.execute(
log_sql, (prefix, data_type, stage_path, phase, status, sql_text, error_msg)
)
def create_table_from_schema(
cursor, data_type: str, stage_path: str, stage_name: str, file_format: str
) -> Optional[str]:
"""
Generate CREATE TABLE SQL using INFER_SCHEMA.
For tables in TABLES_WITH_TRUE_BINARY_DATA, converts VARIANT columns to BINARY
to preserve binary/compressed data.
Args:
cursor: Snowflake cursor for executing queries
data_type: The table name to create
stage_path: The stage path containing the data files
stage_name: The Snowflake stage name
file_format: The Snowflake file format name to use
Returns:
The SQL statement that was generated, or None if no schema could be inferred
"""
# First, locate at least one parquet file to avoid inferring from marker files
# like _SUCCESS that can produce an empty inferred schema.
cursor.execute(f"LIST @{stage_name}/{stage_path}")
list_rows = cursor.fetchall()
parquet_files = [
row[0].rsplit("/", 1)[-1]
for row in list_rows
if row and isinstance(row[0], str) and row[0].lower().endswith(".parquet")
]
if not parquet_files:
return None
sample_file = parquet_files[0]
# Infer schema from a known parquet file.
cursor.execute(f"""
SELECT ARRAY_AGG(OBJECT_CONSTRUCT(*))
FROM TABLE(
INFER_SCHEMA(
LOCATION => '@{stage_name}/{stage_path}',
FILE_FORMAT => '{file_format}',
FILES => ('{sample_file}')
)
)
""")
schema_array = cursor.fetchone()[0]
import json
if isinstance(schema_array, str):
schema_array = json.loads(schema_array)
if not schema_array:
return None
### This was unnessecary when loading prod data ###
# Check if this table should convert BINARY to VARIANT
# convert_binary_to_variant = data_type in TABLES_WITH_VARIANT_COLUMNS
convert_binary_to_variant = False
# Build column definitions manually
column_defs = []
for column_def in schema_array:
col_name = column_def.get("COLUMN_NAME") or column_def.get("column_name")
col_type = column_def.get("TYPE") or column_def.get("type")
nullable = column_def.get("NULLABLE") or column_def.get("nullable", True)
# Skip malformed inference rows that do not define a usable column.
if not col_name or not col_type:
continue
# Convert VARIANT to BINARY for tables with true binary data
if col_type == "VARIANT" and convert_binary_to_variant:
col_type = "BINARY"
# Build column definition
null_clause = "" if nullable else " NOT NULL"
column_defs.append(f'"{col_name}" {col_type}{null_clause}')
if not column_defs:
return None
columns_sql = ",\n ".join(column_defs)
create_sql = f"""
CREATE OR REPLACE TABLE {data_type} (
{columns_sql}
)
"""
return create_sql
def copy_data_into_table(
cursor, data_type: str, stage_path: str, stage_name: str
) -> str:
"""
Generate and execute COPY INTO SQL to load data from stage into table.
Args:
cursor: Snowflake cursor for executing queries
data_type: The table name to copy data into
stage_path: The stage path containing the parquet files
stage_name: The Snowflake stage name
Returns:
The SQL statement that was generated
"""
copy_sql = f"""
COPY INTO {data_type}
FROM @{stage_name}/{stage_path}
FILE_FORMAT = (TYPE = PARQUET BINARY_AS_TEXT = FALSE)
MATCH_BY_COLUMN_NAME = CASE_INSENSITIVE
PATTERN = '^.*\\.parquet$'
"""
return copy_sql
def ensure_masking_policies(cursor, database: str, schema: str):
"""
Idempotently create column masking policies in the schema.
Two policies are created:
- pii_mask_text: masks TEXT/VARCHAR columns
- pii_mask_binary: masks BINARY columns
Both policies return the real value when the session holds the
RAW_TABLE_READ database role, and NULL otherwise.
Args:
cursor: Snowflake cursor for executing queries
database: Database name
schema: Schema name
"""
print("Creating masking policies...")
for policy_name, col_type, return_type in [
("pii_mask_text", "VARCHAR", "VARCHAR"),
("pii_mask_binary", "BINARY", "BINARY"),
]:
cursor.execute(f"""
CREATE MASKING POLICY IF NOT EXISTS {database}.{schema}.{policy_name}
AS (val {col_type}) RETURNS {return_type} ->
CASE WHEN IS_DATABASE_ROLE_IN_SESSION('{DB_ROLE_RAW_TABLE_READ}')
THEN val ELSE NULL END
""")
print(f" ✓ Created masking policy {policy_name}")
print()
def apply_column_masking(cursor, database: str, schema: str, table: str):
"""
Apply masking policies to sensitive columns in a table, if any are defined.
Looks up the table in MASKED_COLUMNS and applies the appropriate policy to
each flagged column. No-op if the table has no masked columns.
Args:
cursor: Snowflake cursor for executing queries
database: Database name
schema: Schema name
table: Table name (case-insensitive; matched against MASKED_COLUMNS keys)
"""
col_policies = MASKED_COLUMNS.get(table.upper(), {})
if not col_policies:
return
for col_name, policy_name in col_policies.items():
cursor.execute(f"""
ALTER TABLE {database}.{schema}.{table}
MODIFY COLUMN {col_name}
SET MASKING POLICY {database}.{schema}.{policy_name}
""")
print(f" ✓ Applied {policy_name} to {table}.{col_name}")
def bootstrap_schema(
cursor,
stack: str,
stage_url: str,
database: str,
stage_name: Optional[str] = None,
):
"""
Bootstrap a new schema for snapshot data loading.
Steps:
1. Create new schema PROD_{stack}
2. Create external stage in the schema
3. Verify stage access
4. Create file format
5. Grant schema privileges to database roles
6. Create masking policies
Args:
cursor: Snowflake cursor for executing queries
stack: Stack identifier (e.g., '576')
stage_url: Complete S3 URL (e.g., 's3://bucket/path/to/data/')
database: Database name
stage_name: Optional custom stage name; if not provided, uses pattern
"""
schema_name = f"PROD_{stack}"
# Determine stage name
if not stage_name:
stage_name = f"prod_{stack}_s3_stage"
# Extract prefix_base from stage_url for verification
# URL format: s3://bucket/path/prefix_base/
prefix_base = stage_url.rstrip("/").split("/")[-1]
print("=" * 70)
print("BOOTSTRAPPING SCHEMA")
print("=" * 70)
print()
# Step 1: Create schema
print(f"Step 1: Creating schema {schema_name}...")
cursor.execute(f"CREATE SCHEMA IF NOT EXISTS {database}.{schema_name}")
cursor.execute(f"USE SCHEMA {database}.{schema_name}")
print(f" ✓ Created schema {schema_name}")
print()
# Step 2: Create external stage
print(f"Step 2: Creating external stage {stage_name}...")
create_stage_sql = f"""
CREATE OR REPLACE STAGE {stage_name}
URL = '{stage_url}'
STORAGE_INTEGRATION = {STORAGE_INTEGRATION}
FILE_FORMAT = (TYPE = PARQUET)
"""
cursor.execute(create_stage_sql)
print(f" ✓ Created stage {stage_name}")
print(f" URL: {stage_url}")
print()
# Step 3: Verify stage access
print(f"Step 3: Verifying stage access...")
test_path = f"{prefix_base}.NODE/1/"
try:
cursor.execute(f"LIST @{stage_name}/{test_path}")
files = cursor.fetchall()
print(
f" ✓ Successfully listed {len(files)} file(s) at @{stage_name}/{test_path}"
)
except Exception as e:
print(f" ⚠ Warning: Could not list files at @{stage_name}/{test_path}")
print(f" Error: {e}")
print(f" This may be expected if the path doesn't exist yet.")
print()
# Step 4: Create file format
print(f"Step 4: Creating file format...")
cursor.execute(f"CREATE OR REPLACE FILE FORMAT {FILE_FORMAT_NAME} TYPE = PARQUET")
print(f" ✓ Created file format {FILE_FORMAT_NAME}")
print()
# Step 5: Grant schema privileges to database roles
print(f"Step 5: Granting schema privileges...")
cursor.execute(
f"GRANT USAGE, MONITOR ON SCHEMA {database}.{schema_name} "
f"TO DATABASE ROLE {database}.{DB_ROLE_RAW_TABLE_READ}"
)
print(f" ✓ Granted USAGE, MONITOR to {database}.{DB_ROLE_RAW_TABLE_READ}")
cursor.execute(
f"GRANT USAGE, MONITOR ON SCHEMA {database}.{schema_name} "
f"TO DATABASE ROLE {database}.{DB_ROLE_RAW_TABLE_READ_CENSORED}"
)
print(f" ✓ Granted USAGE, MONITOR to {database}.{DB_ROLE_RAW_TABLE_READ_CENSORED}")
print()
# Step 6: Create masking policies
print(f"Step 6: Creating masking policies...")
ensure_masking_policies(cursor, database, schema_name)
print("=" * 70)
print("BOOTSTRAP COMPLETE")
print("=" * 70)
print()
return schema_name, stage_name, prefix_base
def grant_table_privileges(cursor, database: str, schema: str):
"""
Grant SELECT and REFERENCES privileges on all tables to database roles.
Args:
cursor: Snowflake cursor for executing queries
database: Database name
schema: Schema name
"""
print("=" * 70)
print("GRANTING TABLE PRIVILEGES")
print("=" * 70)
print()
# Grant to raw_table_read
print(f"Granting SELECT, REFERENCES on all tables to {DB_ROLE_RAW_TABLE_READ}...")
cursor.execute(
f"GRANT SELECT, REFERENCES ON ALL TABLES IN SCHEMA {database}.{schema} "
f"TO DATABASE ROLE {database}.{DB_ROLE_RAW_TABLE_READ}"
)
print(f" ✓ Granted privileges to {database}.{DB_ROLE_RAW_TABLE_READ}")
# Grant to raw_table_read_censored
print(
f"Granting SELECT, REFERENCES on all tables to {DB_ROLE_RAW_TABLE_READ_CENSORED}..."
)
cursor.execute(
f"GRANT SELECT, REFERENCES ON ALL TABLES IN SCHEMA {database}.{schema} "
f"TO DATABASE ROLE {database}.{DB_ROLE_RAW_TABLE_READ_CENSORED}"
)
print(f" ✓ Granted privileges to {database}.{DB_ROLE_RAW_TABLE_READ_CENSORED}")
print()
def setup_temp_tables(cursor):
"""
Create temporary tables needed for the snapshot loading process.
Creates:
- LOAD_LOG: Logs all operations and errors
Args:
cursor: Snowflake cursor for executing queries
"""
print("Setting up logging table...")
# Create LOAD_LOG table
cursor.execute("""
CREATE OR REPLACE TABLE LOAD_LOG (
PREFIX STRING,
DATA_TYPE STRING,
STAGE_PATH STRING,
PHASE STRING, -- e.g. 'START', 'CREATE_TABLE', 'COPY', 'ERROR'
STATUS STRING, -- e.g. 'OK', 'RUN', 'FAILED'
SQL_TEXT STRING, -- SQL we attempted to run (if applicable)
ERROR_MESSAGE STRING, -- populated on failure
LOG_TS TIMESTAMP_NTZ DEFAULT CURRENT_TIMESTAMP()
)
""")
print(" ✓ Created LOAD_LOG table")
def list_prefixes_from_stage(cursor, prefix_base: str, stage_name: str) -> list:
"""
List all prefixes under the given prefix_base from the stage.
Queries the stage to find all directories matching the pattern:
{prefix_base}/{prefix_base}.{data_type}/1/
Args:
cursor: Snowflake cursor for executing queries
prefix_base: Base prefix to search under (e.g. 'dev566')
stage_name: The Snowflake stage name
Returns:
List of prefix strings formatted as '{prefix_base}/{prefix_base}.{data_type}/1/'
"""
print(f"Listing prefixes under {prefix_base}/ in stage...")
# List all files/directories in the stage under prefix_base
cursor.execute(f"""
LIST @{stage_name}/{prefix_base}
""")
results = cursor.fetchall()
prefixes = set()
# Parse the results to extract unique prefix patterns
# Results from LIST contain columns: name, size, md5, last_modified
for row in results:
file_path = row[0] # The 'name' column
# Extract the prefix pattern: {prefix_base}/{prefix_base}.{data_type}/1/
# Example file_path: 's3://synapse-rds-snapshots-dev/test-export/dev566/dev566.NODE/1/part-00000-dcd6d72f-8ee9-400c-94ce-f1f66644c5d3-c000.gz.parquet'
# We want to extract: 'dev566/dev566.NODE/1/'
# The file_path from LIST includes the full S3 path
# Find the prefix_base in the path and extract from there
if f"/{prefix_base}/{prefix_base}." in file_path:
# Find where our prefix pattern starts
idx = file_path.find(f"/{prefix_base}/{prefix_base}.")
# Extract everything after the leading slash
relevant_path = file_path[idx + 1 :]
# Split and reconstruct the prefix pattern
parts = relevant_path.split("/")
if len(parts) >= 3 and parts[2] == "1":
# Reconstruct the prefix: prefix_base/prefix_base.DATA_TYPE/1/
prefix = f"{parts[0]}/{parts[1]}/{parts[2]}/"
prefixes.add(prefix)
prefix_list = sorted(list(prefixes))
print(f" ✓ Found {len(prefix_list)} unique prefixes")
return prefix_list
def process_prefix(
cursor,
prefix: str,
prefix_base: str,
stage_name: str,
file_format: str,
database: str,
schema: str,
) -> Optional[bool]:
"""
Process a single prefix: derive data type, create table, load data,
and apply column masking policies for sensitive tables.
Args:
cursor: Snowflake cursor for executing queries
prefix: The prefix to process (listed from stage)
prefix_base: Base prefix used to derive the stage path and table name
stage_name: The Snowflake stage name
file_format: The Snowflake file format name to use
database: Database name (needed to qualify masking policy references)
schema: Schema name (needed to qualify masking policy references)
Returns:
True if processing succeeded, False if an error occurred
"""
data_type = None
stage_path = None
current_phase = "INIT"
current_sql = None
try:
# Derive data type from prefix
current_phase = "PARSE_PREFIX"
# derive using the provided prefix_base
data_type = derive_data_type(prefix, prefix_base=prefix_base)
# Prefixes are listed as "{prefix_base}/{prefix_base}.{table}/1/", but the
# stage root is already ".../{prefix_base}/" in bootstrap mode. Strip the
# leading "{prefix_base}/" to get a valid stage-relative path.
if prefix.startswith(f"{prefix_base}/"):
stage_path = prefix[len(prefix_base) + 1 :]
else:
stage_path = derive_stage_path(data_type, prefix_base=prefix_base)
infer_location = f"@{stage_name}/{stage_path}"
print(f" Stage location for INFER_SCHEMA: {infer_location}")
# Log: start processing this prefix
log_operation(cursor, prefix, data_type, stage_path, "START", "OK")
# Generate CREATE TABLE SQL
create_sql = create_table_from_schema(
cursor, data_type, stage_path, stage_name, file_format
)
# Some exported table prefixes contain no parquet files (empty table exports).
# In that case, skip table creation/copy without treating it as a failure.
if create_sql is None:
log_operation(
cursor,
prefix,
data_type,
stage_path,
"NO_DATA",
"SKIP",
sql_text=f"INFER_SCHEMA LOCATION => '{infer_location}'",
)
print(
f"Skipping prefix with no inferable schema at {infer_location}",
file=sys.stderr,
)
return None
# Log that we're about to run CREATE TABLE
log_operation(
cursor, prefix, data_type, stage_path, "CREATE_TABLE", "RUN", create_sql
)
# Execute CREATE TABLE
current_phase = "CREATE_TABLE"
current_sql = create_sql
cursor.execute(create_sql)
# Log success
log_operation(cursor, prefix, data_type, stage_path, "CREATE_TABLE", "OK")
# Generate COPY INTO SQL
copy_sql = copy_data_into_table(cursor, data_type, stage_path, stage_name)
print(f" Stage location for COPY INTO: {infer_location}")
# Log that we're about to run COPY INTO
log_operation(
cursor, prefix, data_type, stage_path, "COPY_DATA", "RUN", copy_sql
)
# Execute COPY INTO
current_phase = "COPY_DATA"
current_sql = copy_sql
cursor.execute(copy_sql)
# Log success
log_operation(cursor, prefix, data_type, stage_path, "COPY_DATA", "OK")
# Apply column masking policies for any sensitive columns in this table
current_phase = "APPLY_MASKING"
apply_column_masking(cursor, database, schema, data_type)
return True
except Exception as e:
# Log error with full details
error_msg = f"{type(e).__name__}: {str(e)}"
# Use the tracked phase for accurate error reporting
log_operation(
cursor,
prefix,
data_type if data_type else "UNKNOWN",
stage_path if stage_path else "UNKNOWN",
current_phase,
"FAIL",
sql_text=current_sql,
error_msg=error_msg,
)
print(
f"Error in {current_phase} for prefix {prefix}: {error_msg}",
file=sys.stderr,
)
# Don't re-raise - continue processing other prefixes
return False
def load_snapshot_data(
stage_url: Optional[str] = None,
prefix_base: Optional[str] = None,
stage_name: Optional[str] = None,
database: Optional[str] = None,
schema: Optional[str] = None,
bootstrap_stack: Optional[str] = None,
file_format: Optional[str] = None,
role: Optional[str] = None,
):
"""
Main function to load snapshot data from stage into tables.
Args:
stage_url: Complete S3 URL for the stage (required for bootstrap mode)
prefix_base: Base prefix to search for data (required for manual mode)
stage_name: The Snowflake stage name (required for manual loading)
database: Database name (defaults to SYNAPSE_RDS_SNAPSHOT)
schema: Schema name (required for manual loading)
bootstrap_stack: Stack identifier to bootstrap a new schema (e.g., '576')
file_format: File format name (required for manual mode, auto-created in bootstrap)
role: Snowflake role to use (defaults to SYSADMIN)
"""
# Validate arguments based on mode
if bootstrap_stack:
# Bootstrap mode: require stage_url
if not stage_url:
raise ValueError("--stage-url is required when using --bootstrap-stack")
if not database:
database = DEFAULT_DATABASE
# Extract prefix_base from stage_url in bootstrap mode
prefix_base = stage_url.rstrip("/").split("/")[-1]
# Use default file format name in bootstrap mode
if not file_format:
file_format = FILE_FORMAT_NAME
else:
# Manual loading mode: require schema, stage_name, prefix_base, and file_format
if not stage_name:
raise ValueError(
"--stage-name is required when not using --bootstrap-stack"
)
if not schema:
raise ValueError("--schema is required when not using --bootstrap-stack")
if not prefix_base:
raise ValueError(
"--prefix-base is required when not using --bootstrap-stack"
)
if not file_format:
raise ValueError(
"--file-format is required when not using --bootstrap-stack"
)
if not database:
database = DEFAULT_DATABASE
if not role:
role = "SYSADMIN"
# Connect to Snowflake
conn = snowflake.connector.connect()
cursor = conn.cursor()
try:
# Set role
cursor.execute(f"USE ROLE {role}")
# Set database
cursor.execute(f"USE DATABASE {database}")
# Bootstrap if requested
if bootstrap_stack:
# stage_url is guaranteed to be set by validation above
assert stage_url is not None
schema_name, stage_name, prefix_base = bootstrap_schema(
cursor=cursor,
stack=bootstrap_stack,
stage_url=stage_url,
database=database,
stage_name=stage_name,
)
schema = schema_name
conn.commit()
# Set schema (required at this point)
cursor.execute(f"USE SCHEMA {schema}")
# Set up temporary tables (LOAD_LOG)
setup_temp_tables(cursor)
conn.commit()
# Ensure prefix_base, stage_name, and file_format are set by this point
# (guaranteed by validation logic above)
assert prefix_base is not None, "prefix_base should be set by this point"
assert stage_name is not None, "stage_name should be set by this point"
assert file_format is not None, "file_format should be set by this point"
# List all prefixes from the stage under prefix_base
prefixes = list_prefixes_from_stage(cursor, prefix_base, stage_name)
print(f"Processing {len(prefixes)} prefixes...")
# Track successes and failures
success_count = 0
failure_count = 0
# Process each prefix
for prefix in prefixes:
print(f"Processing prefix: {prefix}")
result = process_prefix(
cursor, prefix, prefix_base, stage_name, file_format,
database=database, schema=schema,
)
conn.commit() # Commit after each prefix
if result is True:
success_count += 1
elif result is False:
failure_count += 1
print(f"\n{'='*60}")
print("Processing complete!")
print(f" ✓ Successful: {success_count}")
print(f" ✗ Failed: {failure_count}")
print(f" Total: {len(prefixes)}")
print(f"{'='*60}")
# Grant table privileges if bootstrapped
if bootstrap_stack and schema:
grant_table_privileges(cursor, database, schema)
conn.commit()
except Exception as e:
print(f"Fatal error: {e}", file=sys.stderr)
conn.rollback()
raise
finally:
cursor.close()
conn.close()
if __name__ == "__main__":
import argparse
parser = argparse.ArgumentParser(
description=(
"Load RDS snapshot data from S3 into Snowflake tables. "
"Two modes: (1) Bootstrap mode - use --bootstrap-stack + --stage-url to create "
"schema, stage, and load data. (2) Manual mode - use --schema + --stage-name + "
"--prefix-base + --file-format to load data into existing schema/stage."
)
)
parser.add_argument(
"--bootstrap-stack",
dest="bootstrap_stack",
default=None,
help=(
"Stack number (e.g., '576') to bootstrap a new schema PROD_{stack}. "
"Creates schema, external stage, file format, and grants privileges. "
"Requires --stage-url."
),
)
parser.add_argument(
"--stage-url",
dest="stage_url",
default=None,
help=(
"Complete S3 URL for the stage location (e.g., "
"'s3://synapse-rds-snapshots-dev/prod-576-export/dev123/'). "
"Required when using --bootstrap-stack."
),
)
parser.add_argument(
"--prefix-base",
dest="prefix_base",
default=None,
help=(
"Base prefix to search for data in the stage (e.g., 'dev123'). "
"Required when NOT using --bootstrap-stack (manual mode)."
),
)
parser.add_argument(
"--file-format",
dest="file_format",
default=None,
help=(
f"File format name (default in bootstrap mode: {FILE_FORMAT_NAME}). "
"Required when NOT using --bootstrap-stack (manual mode)."
),
)
parser.add_argument(
"--database",
dest="database",
default=DEFAULT_DATABASE,
help=f"Database name (default: {DEFAULT_DATABASE})",
)
parser.add_argument(
"--schema",
dest="schema",
default=None,
help=(
"Schema name. Required when NOT using --bootstrap-stack (manual mode). "
"Auto-generated as PROD_{stack} in bootstrap mode."
),
)
parser.add_argument(
"--stage-name",
dest="stage_name",
default=None,
help=(
"Snowflake stage name. Required when NOT using --bootstrap-stack (manual mode). "
"Auto-generated as prod_{stack}_s3_stage in bootstrap mode."
),
)
parser.add_argument(
"--role",
dest="role",
default="SYSADMIN",
help="Snowflake role to use (default: SYSADMIN)",
)
args = parser.parse_args()
load_snapshot_data(
stage_url=args.stage_url,
prefix_base=args.prefix_base,
stage_name=args.stage_name,
database=args.database,
schema=args.schema,
bootstrap_stack=args.bootstrap_stack,
file_format=args.file_format,
role=args.role,
)
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment