import os import logging from time import sleep from google.cloud.storage import Client from google.oauth2.service_account import Credentials from vector_utils.connections.connection import Connection from vector_utils.connections import exceptions class GCSConnection(Connection): """The connection class is used for GCS connections. Attributes: conn_obj (ConnectionInfo) logger (obj): Provide a logging object """ def __init__(self, conn_obj, logger=None): super().__init__(conn_obj) self.connection = None self._logger = logger or self._setup_logger() self.login() @property def logger(self): return self._logger @logger.setter def logger(self, value): self._logger = value def _setup_logger(self): """Setup logger helper.""" logger = logging.getLogger(__name__) logger.addHandler(logging.NullHandler()) return logger def login(self) -> None: """Login and setup connection to GCS.""" credentials = Credentials.from_service_account_info(self.conn_obj.gcs_config_file) self._gcs_client = Client(credentials=credentials, project=credentials.project_id) self._bucket = self._gcs_client.bucket(self.conn_obj.domain_name) @property def bucket(self): """Always return a valid bucket.""" return self._bucket def _normalize_path(self, path, is_dir=False): """Normalize blob path by stripping leading '/' and ensuring trailing '/' for dirs.""" normalized = path.lstrip('/') if is_dir and normalized and not normalized.endswith('/'): normalized += '/' return normalized @Connection.handle_exceptions def check_connection(self, remote_initial_dir): pass @Connection.handle_exceptions def close_connection(self): """Close the client connection.""" pass @Connection.handle_exceptions def delete(self, filename, remove_dir=False): pass @Connection.handle_exceptions def file_exists(self, file_name): """Check file exists. Args: file_name (str): A valid filename path Returns: bool """ return self.bucket.blob(self._normalize_path(file_name)).exists() @Connection.handle_exceptions def file_size(self, file_name): """Return file size. Args: file_name (str): Full file path Returns: int: The size in bytes. """ blob = self.bucket.get_blob(self._normalize_path(file_name)) if blob is None: raise FileNotFoundError('No such file') return blob.size @Connection.handle_exceptions def is_dir(self, dirname): """Change directory. Args: dirname (str): The directory to change to. """ pass @Connection.handle_exceptions def mkdir(self, dest_dir: str) -> bool | None: """Create directories recursively. Args: dest_dir (str): The full directory path. Returns: bool: True if we were able to create the path """ dest_dir = self._normalize_path(dest_dir, is_dir=True) dest_dir_blob = self.bucket.blob(dest_dir) if not dest_dir_blob.exists(): dest_dir_blob.upload_from_string(b'') return True return None @Connection.handle_exceptions def rmdir(self, dir_name): """Remove the directory path. Args: dir_name (str): The full directory path. """ pass @Connection.handle_exceptions def reconnect(self): """Reconnect connection.""" self.login() @Connection.handle_exceptions def scan_dir(self, directory, exceptions=[]): """Return list of directories or files. Args: directory (str): The path which you which to scan. exceptions (list): The files or folders to exclude. Returns: list(str) """ directory = self._normalize_path(directory, is_dir=True) results = self.bucket.list_blobs(prefix=directory) directory_results = [] for blob in results: name = os.path.relpath(blob.name, directory) if name and name not in exceptions: directory_results.append(blob.name) return directory_results @Connection.handle_exceptions def transfer_files(self, file_list, transfer_mode='upload', batch_file_id=None, pid=None, overwrite=None, check_delivered_files=True, d_job=None): """Transfer files up or down. Args: file_list (list(dict)): A list of remote/local keypairs transfer_mode (str): upload or download batch_file_id: pid: overwrite: check_delivered_files: d_job: Returns: Raises: Exception: We are unable to transfer files. """ attempts = 1 while attempts <= 5: self.logger.info('Delivery attempt # {}.'.format(attempts)) try: for file in file_list: local_file = file['local'] remote_file = file['remote'] self.logger.info('Start transferring.') blob = self.bucket.blob(self._normalize_path(remote_file)) if transfer_mode == 'download': with open(local_file, 'wb') as data: blob.download_to_file(data) else: with open(local_file, 'rb') as data: blob.upload_from_file(data) self.logger.info('Transfer finished.') self.logger.info('Delivery succeeded.') return True except Exception as e: self.logger.error('Error transferring file: {}'.format(e)) attempts += 1 sleep(15) self.logger.info('Delivery failed.') raise exceptions.TransferException()