"""SFTP Connection Object Wrapper.""" import io import os from typing import List from typing import Tuple import paramiko def tidy(handler): """Wrap function, close connction on error, handle cwd.""" def wrapper(self, *args, **kwargs): start_pointer = self._sftp_client.getcwd() try: results = handler(self, *args, **kwargs) except Exception as e: self.close() raise e else: self._sftp_client.chdir(start_pointer) return results return wrapper class Connection(): # noqa:D101 _sftp_client: paramiko.sftp_client.SFTPClient _ssh_client: paramiko.SSHClient def __init__(self, hostname: str, port: int, username: str, pkey: str): """Initialize SFTP connection. Args: hostname (str): target host to connct port (int): target port to connect on host username (str): auth username on host pkey (str): auth private key on host """ self._ssh_client = paramiko.SSHClient() self._ssh_client.set_missing_host_key_policy(paramiko.AutoAddPolicy()) # noqa:E501 self._ssh_client.connect( hostname=hostname, port=port, username=username, pkey=paramiko.RSAKey.from_private_key(io.StringIO(pkey)) ) self._sftp_client = self._ssh_client.open_sftp() # type: ignore def close(self) -> None: """Terminate sftp connection.""" self._ssh_client.close() @tidy def upload_files(self, file_list: List[Tuple[bytes, str]]) -> None: """Upload file(s). Args: file_list list(tuple): - local (bytes): data in memory to transfer - remote (str): path on server to upload to Returns: None """ directories = sorted( {os.path.dirname(x[1]) for x in file_list}, reverse=True ) for directory in directories: self.mkdir(directory) for (source, destination) in file_list: self._sftp_client.putfo( io.BytesIO(source), destination ) @tidy def mkdir(self, new_dir: str) -> bool: """Create directories in path. Args: new_dir (str): The full directory path. Returns: bool: if directory was created """ created = False if new_dir in ('/', ''): return created try: parts = new_dir.strip('/').split('/') for part in parts: # move up dir, create if not exists try: self._sftp_client.chdir(part) except (IOError, FileNotFoundError): # create dir, OSError if dir already exists try: self._sftp_client.mkdir(part) except OSError: pass else: created = True # move up dir, verifying new dir existance self._sftp_client.chdir(part) except Exception as e: raise e return created @tidy def scan_dir(self, directory: str) -> List[str]: """List files/directories. Args: directory (str): folder to search Returns: list(str): file and/or directory names """ results = self._sftp_client.listdir(path=directory) return [x for x in results]