from __future__ import annotations import logging from collections.abc import Iterator from dataclasses import dataclass from typing import Any import snowflake.connector from snowflake.connector import DictCursor from snowflake.connector.connection import SnowflakeConnection from snowflake.connector.converter import SnowflakeConverter from app.connectors.database.sql import GET_DOMAINS_TO_VERIFY, UPDATE_DOMAINS from app.connectors.database.types import EmailDomainValidation logger = logging.getLogger(__name__) @dataclass class SnowflakeClient: _conn: SnowflakeConnection | None = None def __init__( self, env: str, user: str, account: str, warehouse: str, database: str, role: str, schema: str, private_key: bytes, ) -> None: self.env = env self.user = user self.account = account self.warehouse = warehouse self.database = database self.role = role self.schema = schema self.private_key = private_key @property def conn(self) -> SnowflakeConnection: if self._conn is None: self._conn = snowflake.connector.connect( user=self.user, private_key=self.private_key, account=self.account, warehouse=self.warehouse, database=self.database, role=self.role, schema=self.schema, ) return self._conn def get_domains_to_verify( self, domains_chunk_size: int = 100 ) -> Iterator[list[EmailDomainValidation]]: with self.conn.cursor(DictCursor) as cur: cur.execute(GET_DOMAINS_TO_VERIFY) while domains_chunk := cur.fetchmany(domains_chunk_size): yield [ EmailDomainValidation.model_validate(row) for row in domains_chunk ] def update_domains(self, domains: list[EmailDomainValidation]) -> None: logger.info("Updating db") dumped_rows = tuple(tuple(row.model_dump().values()) for row in domains) escaped_data = self._escape_values(dumped_rows) # Prevent leading comma in data tuple escaped_data = escaped_data[0] if len(escaped_data) == 1 else escaped_data # Here we are filling template query with actual data # and removing outer parentheses to follow snowflake/sql syntax # also we are replacing python's None objests with null for snowflake/sql values = ( f"{escaped_data}".replace("((", "(") .replace("))", ")") .replace("None", "null") ) query = UPDATE_DOMAINS.format(values=values) with self.conn.cursor() as cur: cur.execute(query) @classmethod def _escape_values(cls, data: Any) -> Any: if isinstance(data, str): return SnowflakeConverter.escape(data) if isinstance(data, list | tuple): return type(data)(cls._escape_values(el) for el in data) return data def close(self) -> None: if self._conn: self._conn.close()