"""Database connection and query execution utilities.""" from contextlib import contextmanager from typing import 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) -> pd.DataFrame: """ Execute a SQL query and return results as DataFrame. Args: session: Active Snowpark session query: SQL query string (already formatted with values) Returns: Pandas DataFrame with query results """ result = session.sql(query) return result.to_pandas() def execute_non_query(session: Session, query: str) -> None: """ Execute a non-query SQL statement (INSERT, UPDATE, DELETE, MERGE). Args: session: Active Snowpark session query: SQL statement string (already formatted with values) """ session.sql(query).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 def get_current_user(session: Session) -> str: """ Get the current Snowflake user and role running the application. For Snowflake Streamlit in Snowflake apps, uses st.experimental_user. For local development, uses CURRENT_USER() and CURRENT_ROLE() SQL functions. Args: session: Active Snowpark session Returns: User and role string (e.g., "OZHOVNUVATYI - FANSIFTER_ENGINEERING") """ # Cache in session state to avoid repeated queries if "current_user" not in st.session_state: try: # Try Streamlit in Snowflake user context first if hasattr(st, "experimental_user"): user_info = st.experimental_user # experimental_user returns dict with 'user_name' key if isinstance(user_info, dict) and "user_name" in user_info: username = user_info["user_name"] # Get role separately result = session.sql("SELECT CURRENT_ROLE() as role").collect() role = result[0]["ROLE"] st.session_state.current_user = f"{username} - {role}" else: # Fallback to SQL query result = session.sql( "SELECT CURRENT_USER() as username, CURRENT_ROLE() as role" ).collect() username = result[0]["USERNAME"] role = result[0]["ROLE"] st.session_state.current_user = f"{username} - {role}" else: # Fallback to SQL query for local development result = session.sql( "SELECT CURRENT_USER() as username, CURRENT_ROLE() as role" ).collect() username = result[0]["USERNAME"] role = result[0]["ROLE"] st.session_state.current_user = f"{username} - {role}" except Exception: # Final fallback - use SQL query result = session.sql( "SELECT CURRENT_USER() as username, CURRENT_ROLE() as role" ).collect() username = result[0]["USERNAME"] role = result[0]["ROLE"] st.session_state.current_user = f"{username} - {role}" return st.session_state.current_user