"""Database connection and query execution utilities.""" from contextlib import contextmanager from typing import Any, Iterator import pandas as pd import streamlit as st from snowflake.snowpark.context import get_active_session from snowflake.snowpark.exceptions import SnowparkSessionException from snowflake.snowpark.session import Session from common.local_connection import get_connection_parameters @st.cache_resource(show_spinner="Connecting to Snowflake...") def get_session() -> Session: """ Get or create a Snowflake Snowpark session. Returns cached session for performance. In Snowflake deployment, uses active session. For local development, creates session from connections.json. Returns: Active Snowpark session """ try: # Try to get active session (Snowflake deployment) session = get_active_session() except SnowparkSessionException: # Local development - create session from connections file connection_params = get_connection_parameters("artist_roster") session = Session.builder.configs(connection_params).create() return session def execute_query( session: Session, query: str, params: dict[str, Any] | None = None ) -> pd.DataFrame: """ Execute a parameterized SQL query and return results as DataFrame. Args: session: Active Snowpark session query: SQL query string with :param_name placeholders params: Dictionary of parameter names to values Returns: Pandas DataFrame with query results """ if params is None: params = {} # Find all parameter names in the query in order of appearance import re param_names = re.findall(r":(\w+)", query) # Create ordered list of parameter values based on query order param_list = [params.get(name) for name in param_names] # Replace named parameters with ? placeholders in order query_with_placeholders = re.sub(r":\w+", "?", query) # Execute query with parameters as list result = session.sql(query_with_placeholders, params=param_list) return result.to_pandas() def execute_non_query( session: Session, query: str, params: dict[str, Any] | None = None ) -> None: """ Execute a non-query SQL statement (INSERT, UPDATE, DELETE). Args: session: Active Snowpark session query: SQL statement string with :param_name placeholders params: Dictionary of parameter names to values """ if params is None: params = {} # Find all parameter names in the query in order of appearance import re param_names = re.findall(r":(\w+)", query) # Create ordered list of parameter values based on query order param_list = [params.get(name) for name in param_names] # Replace named parameters with ? placeholders in order query_with_placeholders = re.sub(r":\w+", "?", query) session.sql(query_with_placeholders, params=param_list).collect() @contextmanager def transaction(session: Session) -> Iterator[Session]: """ Context manager for database transactions. Usage: with transaction(session) as tx: execute_non_query(tx, "INSERT ...", params) Args: session: Active Snowpark session Yields: Session within transaction context """ try: # Snowpark handles transactions automatically # This is a placeholder for explicit transaction control if needed yield session except Exception as e: # In case of error, Snowpark will rollback automatically raise e