from typing import Dict, List, Tuple import pandas as pd import sqlparse from snowflake.snowpark.session import Session def get_snowflake_session( snowflake_configs: dict, warehouse: str, database: str, schema: str ) -> Session: return Session.builder.configs( {"warehouse": warehouse, "database": database, "schema": schema, **snowflake_configs} ).create() def get_sql_result(session: Session, schema: str): df = session.sql(schema) local_df = df.collect() if len(local_df) == 0: raise Exception("No data found") return pd.DataFrame(local_df) def extract_columns_and_data_types_from_sql( sql_text: str, required_dtypes: List[str] = ["number", "date", "varchar(255)"] ) -> List[Tuple[str, str]]: """ Parse DDL to get column names and their exact SQL data types. Returns a list of (column_name, sql_data_type) tuples in order of definition. The parsing is done by selecting the SQL token that is parenthesized. To avoid incorrectly selecting tokens such as "CLUSTER BY (p_date)", we only select parenthesized tokens that contain all of the specified data types. These can be customized for tables with different data types. """ # Parse SQL into tokens parsed = sqlparse.parse(sql_text)[0] columns = [] for token in parsed.tokens: # Find parenthesized tokens that contain all of the specified data types if isinstance(token, sqlparse.sql.Parenthesis) and all( [dtype in token.value.lower() for dtype in required_dtypes] ): column_defs = token.value.strip("()").split(",") for col_def in column_defs: col_def = col_def.strip() # Remove empty lines and SQL comments if not col_def or col_def.startswith("--"): continue # Split by whitespace for initial parsing parts = col_def.split() if len(parts) < 2: continue # Extract column name column_name = parts[0] # Handle inline comments with COMMENT keyword and inline SQL comments remaining = col_def[len(column_name) :].strip() data_type = remaining.split("COMMENT")[0].split("--")[0].strip() columns.append((column_name, data_type)) if columns: return columns else: raise Exception("No columns found in the parsed SQL: \n" + sql_text) def create_schedule_task(session: Session, task_name: str, warehouse: str, schedule: str, func: str): schema = """ CREATE OR REPLACE TASK {task_name} WAREHOUSE = {warehouse} SCHEDULE = 'USING CRON {schedule} UTC' AS DECLARE p_date string; begin p_date := DATE(SYSDATE() - INTERVAL '1 DAY'); CALL {func}(:p_date); end; """.format(task_name=task_name, warehouse=warehouse, schedule=schedule, func=func) get_sql_result(session, schema) def get_table_names(session: Session, database: str, schema: str, table_prefix: str = "") -> List[str]: """Get all table names from a given database and schema""" query = f""" SELECT table_name FROM {database}.INFORMATION_SCHEMA.TABLES WHERE table_schema = '{schema}' AND table_name LIKE '{table_prefix}%' """ result = session.sql(query).collect() return sorted([row["TABLE_NAME"] for row in result]) def get_table_columns( session: Session, database: str, schema: str, table_names: List[str] ) -> Dict[str, List[str]]: """Get all table columns from a given database and schema""" table_names_quoted = [f"'{table_name}'" for table_name in table_names] query = f""" SELECT table_name, column_name FROM {database}.INFORMATION_SCHEMA.COLUMNS WHERE 1=1 AND table_schema = '{schema}' AND table_name IN ({", ".join(table_names_quoted)}) ORDER BY table_name, ordinal_position """ result = session.sql(query).collect() table_columns = {} for row in result: table_name = row["TABLE_NAME"] column_name = row["COLUMN_NAME"] if table_name not in table_columns: table_columns[table_name] = [] table_columns[table_name].append(column_name) return table_columns