from typing import Any, Dict, List, Optional import backoff from snowflake import connector as sf_conn from snowflake.connector.cursor import DictCursor from snowflake.connector.errors import DatabaseError, OperationalError from config import SNOWFLAKE_ACCOUNT, SNOWFLAKE_DATABASE, SNOWFLAKE_HOST, SNOWFLAKE_PASSWORD, SNOWFLAKE_PORT, \ SNOWFLAKE_SCHEMA, SNOWFLAKE_USER, SNOWFLAKE_WAREHOUSE from constants import BACKOFF_TIMEOUT from .db_connector import DBConnector __all__ = ["SnowflakeConnector", "SnowflakeConnectionError"] class SnowflakeConnectionError(Exception): pass class SnowflakeConnector(DBConnector): user = SNOWFLAKE_USER password = SNOWFLAKE_PASSWORD host = SNOWFLAKE_HOST port = SNOWFLAKE_PORT database = SNOWFLAKE_DATABASE warehouse = SNOWFLAKE_WAREHOUSE account = SNOWFLAKE_ACCOUNT schema = SNOWFLAKE_SCHEMA def close(self): self._conn.close() def _get_connection(self): try: return sf_conn.connect( user=self.user, password=self.password, host=self.host, port=self.port, database=self.database, warehouse=self.warehouse, account=self.account, schema=self.schema, timezone="UTC", ) except DatabaseError as e: raise SnowflakeConnectionError("Can't connect to snowflake DB") from e @backoff.on_exception(backoff.expo, OperationalError, max_time=BACKOFF_TIMEOUT) def execute_query(self, query: str, params: Optional[Dict[str, Any]] = None) -> List[Dict[str, Any]]: """ Execute snowflake query :param query: Raw SQL string :param params: SQL query parameters :return: List of dicts with uppercase column names in keys """ with self._conn.cursor(DictCursor) as cursor: result = cursor.execute(query, params, _no_retry=True).fetchall() return result @backoff.on_exception(backoff.expo, OperationalError, max_time=BACKOFF_TIMEOUT) def execute_async_query(self, query: str, params: Optional[Dict[str, Any]] = None) -> str: with self._conn.cursor(DictCursor) as cursor: query_id = cursor.execute_async(query, params).get("queryId") return query_id @backoff.on_exception(backoff.expo, OperationalError, max_time=BACKOFF_TIMEOUT) def get_query_result(self, query_id) -> List[Dict[str, Any]]: with self._conn.cursor(DictCursor) as cursor: result = cursor.query_result(query_id).fetchall() return result def check_db_connection(self) -> bool: results = self.execute_query("SELECT 1") return results == [{"1": 1}]