import json import os from typing import Any, Dict, Optional import boto3 import pandas as pd import snowflake.connector from snowflake.connector.pandas_tools import write_pandas from .snowflake.constants import Warehouse, Role from cryptography.hazmat.primitives import serialization from cryptography.hazmat.backends import default_backend from dotenv import load_dotenv def get_snowflake_private_key(): raw_private_key = os.environ["SNOWFLAKE_PRIVATE_KEY"] # Handle newlines: replace literal '\n' with actual newlines if "\\n" in raw_private_key: raw_private_key = raw_private_key.replace("\\n", "\n") try: private_key = serialization.load_pem_private_key( raw_private_key.encode("utf-8"), password=None, backend=default_backend() ) return private_key except ValueError as e: print(f"Failed to load private key: {str(e)}") raise def get_snowflake_connection( warehouse: Optional[Warehouse] = Warehouse.SMALL, role: Optional[Role] = None ): """Create and return a Snowflake connection. Args: warehouse: Snowflake warehouse to use (default: SUNO_PROD_X_SMAL) role: Snowflake role to use (default: None, uses user's default role) """ connection_params = { "user": os.environ["SNOWFLAKE_ACCOUNT_USER"], "private_key": get_snowflake_private_key(), "account": os.environ["SNOWFLAKE_ACCOUNT"], "warehouse": warehouse.value if warehouse else Warehouse.SMALL.value, "database": "SUNO_PROD", "schema": "PROD", } if role: connection_params["role"] = role.value return snowflake.connector.connect(**connection_params) def execute_snowflake_ddl( ddl_statement: str, warehouse: str = Warehouse.LARGE, # default to a larger warehouse for DDL operations role: Optional[Role] = Role.ACCOUNTADMIN, ) -> None: """Execute a DDL statement against Snowflake. Args: ddl_statement: DDL statement to execute (CREATE, ALTER, DROP, etc.) warehouse: Snowflake warehouse to use (default: SUNO_PROD_LARGE) role: Snowflake role to use (default: ACCOUNTADMIN) """ with get_snowflake_connection(warehouse=warehouse, role=role) as conn: try: cursor = conn.cursor() cursor.execute(ddl_statement) except Exception as e: raise Exception(f"Error executing Snowflake DDL: {e!s}") def query_snowflake( query: str, params: Optional[Dict[str, Any]] = None, warehouse: Optional[Warehouse] = Warehouse.SMALL, role: Optional[Role] = None, ) -> pd.DataFrame: """Execute a query against Snowflake and return results as a DataFrame. Args: query: SQL query string to execute params: Optional dictionary of query parameters warehouse: Snowflake warehouse to use (default: SUNO_PROD_X_SMAL) role: Snowflake role to use (default: None, uses user's default role) Returns: DataFrame containing query results """ with get_snowflake_connection(warehouse=warehouse, role=role) as conn: try: cursor = conn.cursor() if params: cursor.execute(query, params) else: cursor.execute(query) results = cursor.fetchall() columns = [desc[0] for desc in cursor.description] return pd.DataFrame(results, columns=columns) except Exception as e: raise Exception(f"Error executing Snowflake query: {e!s}") def write_to_snowflake( df: pd.DataFrame, table_name: str, schema: Optional[str] = None, warehouse: Optional[Warehouse] = Warehouse.SMALL, role: Optional[Role] = None, ) -> None: """Write a DataFrame to a Snowflake table. Args: df: DataFrame to write table_name: Target table name schema: Optional schema name (defaults to connection schema) warehouse: Snowflake warehouse to use (default: SUNO_PROD_X_SMAL) role: Snowflake role to use (default: None, uses user's default role) """ with get_snowflake_connection(warehouse=warehouse, role=role) as conn: try: return write_pandas(conn, df, table_name, schema=schema, on_error="continue") except Exception as e: raise Exception(f"Error writing to Snowflake: {e!s}") def get_lambda_client(): """Create and return an AWS Lambda client.""" return boto3.client( "lambda", aws_access_key_id=os.environ["AWS_ACCESS_KEY_ID"], aws_secret_access_key=os.environ["AWS_SECRET_ACCESS_KEY"], region_name="us-east-2", ) def get_s3_client(): """Create and return an AWS S3 client.""" return boto3.client( "s3", aws_access_key_id=os.environ["AWS_ACCESS_KEY_ID"], aws_secret_access_key=os.environ["AWS_SECRET_ACCESS_KEY"], region_name="us-east-2", ) def invoke_lambda(function_name: str, payload: Dict[str, Any]) -> Dict[str, Any]: """Invoke an AWS Lambda function and return the response.""" lambda_client = get_lambda_client() response = lambda_client.invoke( FunctionName=function_name, InvocationType="RequestResponse", Payload=json.dumps(payload), ) return json.loads(response["Payload"].read().decode("utf-8")) def invoke_postgres_lambda(query, values=None, is_reader=True): """Invoke the Postgres Lambda, which proxies queries from Dagster.""" payload = {"query": query, "values": values, "is_reader": is_reader} return invoke_lambda( function_name="database-accessor", payload=payload, ) def invoke_redis_lambda( operation: str, keys: list[str], values: Optional[list[dict]] = None, ttl: Optional[int] = None ) -> Any: """Invoke the Redis Lambda function to perform Redis operations. Args: operation: The Redis operation to perform ('get', 'set_json', or 'delete') keys: A list of keys to operate on values: The values to set. This should be a list of values matching the length of keys. ttl: Optional time-to-live in seconds (only used for 'set' operation) """ valid_operations = ["get", "set_json", "delete"] operation = operation.lower() if operation not in valid_operations: raise ValueError(f"Invalid operation: {operation}. Must be one of {valid_operations}") # Validate values for set operation if operation == "set_json": if values is None: raise ValueError("Values must be provided for 'set_json' operation") if len(values) != len(keys): raise ValueError( f"Number of values ({len(values)}) must match number of keys ({len(keys)})" ) payload = { "operation": operation, "keys": keys, } if operation == "set_json": payload["values"] = values if ttl is not None: payload["ttl"] = ttl return invoke_lambda( function_name="redis-generic-accessor", payload=payload, )