diff --git a/docs/changes.rst b/docs/changes.rst index 7f6936f7..6d6bc37b 100644 --- a/docs/changes.rst +++ b/docs/changes.rst @@ -1,5 +1,12 @@ -1.2.3 (current, released 2026-9-08) +1.2.4 (current, released 2026-9-24) ----------------------------------- + * adding new tests for localtree and retry helpers. + * adding key_type parameter to get_hostkey to pull specific key. + * change hostkey verification comparison to pass bytes to hash function. + * fix remote path nesting and duplicate sub-directory names in localtree. + +1.2.3 (released 2026-9-08) +-------------------------- * adding new tests for connections, hash, drivepath and read-only servers. * reworking tests to incorporate self-cleaning of all artifacts. * fix for error handling catches in _set_authentication function. diff --git a/docs/conf.py b/docs/conf.py index d9d882cd..82fe0410 100644 --- a/docs/conf.py +++ b/docs/conf.py @@ -54,9 +54,9 @@ # built documents. # # The short X.Y version. -version = '1.2.3' +version = '1.2.4' # The full version, including alpha/beta/rc tags. -release = '1.2.3' +release = '1.2.4' # The language for content autogenerated by Sphinx. Refer to documentation # for a list of supported languages. diff --git a/docs/contributing.rst b/docs/contributing.rst index c805c1fb..e6aa305d 100644 --- a/docs/contributing.rst +++ b/docs/contributing.rst @@ -21,7 +21,7 @@ Code a. Setup CI testing for your fork. Currently testing is done on Github Actions but feel free to use the framework of your choosing. b. Testing features that concern chmod, chown on Windows is NOT supported. Testing compression has to be ran against a local compatible sshd and not the pytest-sftpserver plugin as it does NOT support this feature. - c. You will need to setup an ssh daemon on your local machine and create a user: copy the contents of id_sftpretty.pub to the newly created user's authorized_keys file -- Tests that can only be ran locally are skipped using the @skip_if_ci decorator so they don't fail when the test suite runs on the CI server. + c. You will need to setup an ssh daemon on your local machine and create a user. Copy the contents of id_sftpretty.pub to the newly created user's authorized_keys file. Tests that can only be ran locally are skipped using the @SKIP_IF_CI decorator so they don't fail when the test suite runs on the CI server. #. Ensure that your name is added to the end of the :doc:`authors` file using the format Name (url), where the (url) portion is optional. #. Submit a Pull Request to the project. @@ -36,9 +36,9 @@ This section lists the priority that will be assigned to an issue: #. Developer Issues #. Issues that have a pull request with a test(s) displaying the issue and code change(s) that satisfies the test suite #. Issues that have a pull request with a test(s) displaying the issue - #. Naked pull requests - a code change request with no accompaning test + #. Naked pull requests. A code change request with no accompaning test #. An issue without a pull request with a test displaying the issue - #. Badly documented issue with no code or test - sftpretty is not an end-user tool, it is a developer tool and it is expected that issues will be submitted like a developer and not an end-user. Issues in the realm of "the internet is broken" will be marked as invalid with a comment pointing the submitter to this section. + #. Badly documented issue with no code or test. sftpretty is not an end-user tool, it is a developer tool and it is expected that issues will be submitted like a developer and not an end-user. Issues in the realm of "the internet is broken" will be marked as invalid with a comment pointing the submitter to this section. Testing ------- diff --git a/docs/cookbook.rst b/docs/cookbook.rst index 3425f96d..b067eb12 100644 --- a/docs/cookbook.rst +++ b/docs/cookbook.rst @@ -532,37 +532,145 @@ Don't like how we have modified a paramiko method? Use this attribute to get at the original version. Our goal is to augment not supplant paramiko. -:func:`sftpretty.localtree` ---------------------------- +:func:`sftpretty.helpers._callback` +----------------------------------- +A progress reporter implementation :meth:`sftpretty.Connection.get` and +:meth:`sftpretty.Connection.put` reach for when you hand them no ``callback`` +of your own. Paramiko calls it once per chunk with the bytes moved so far and +the total. A large file can fill your terminal, give it a ``logger`` and let +your handler save the transfer output to a preferred location. Roll your own +callable that takes two integers if you dare. + +.. code-block:: python + + >>> from logging import getLogger + >>> from sftpretty.helpers import _callback + + >>> _callback('eels.txt', 512, 1024) + Transfer of File: [eels.txt] @ 50.0% 512:1024 bytes + + >>> log = getLogger('LoggyMcLogs') + >>> _callback('eels.txt', 512, 1024, logger=log) + + +:func:`sftpretty.helpers.drivepath` +----------------------------------- +A painful shim that attempts to convert Windows based pathing into a valid +POSIX one, that the remote will accept. It's purely lexical, nothing is opened, +resolved or checked for existence. A path already in POSIX form returns +untouched. + +.. code-block:: python + + >>> from sftpretty.helpers import drivepath + + >>> drivepath('C:\\Users\\nick\\file.txt') + '/C:/Users/nick/file.txt' + >>> drivepath('C:tmp\\test.txt') + '/C:/tmp/test.txt' + >>> drivepath('\\\\server\\share\\file.txt') + '//server/share/file.txt' + >>> drivepath('/home/user/file.txt') + '/home/user/file.txt' + + +:func:`sftpretty.helpers.hash` +------------------------------ + +One digest, five types of input. Give it a path, an open file object, a +:class:`io.BytesIO`, some bytes or a string and get back the hexdigest. Anything +other than the five input types digest as an empty buffer rather than +complianing. A string that cannot be opened is digested as text rather than +raising. Only the ``algorithm.name`` is read, so any spent hash object can be +passed without remnants carrying over between calls. Files are read in +``blocksize`` chunks, so size shouldn't be a concern. + +.. code-block:: python + + >>> from hashlib import md5 + >>> from pathlib import Path + >>> from sftpretty.helpers import hash + + >>> Path('/tmp/eels.txt').write_text('My hovercraft is full of eels.') + 30 + >>> hash('/tmp/eels.txt') == hash('My hovercraft is full of eels.') + True + >>> hash(open('/tmp/eels.txt', 'rb')) == hash('/tmp/eels.txt') + True + >>> hash('/tmp/eels.txt', algorithm=md5()) + '5d5bc914f200b729e1c64c927cafe8c3' + >>> hash('/tmp/eels.txt', blocksize=8192) + '4953167ab20a15c0...' + + +:func:`sftpretty.helpers.localtree` +----------------------------------- Similar to :meth:`sftpretty.Connection.remotetree` except that it walks a **local** directory structure. It has the same output format and likewise -stores the resulting tree in a dictionary. +stores the resulting tree in a dictionary. Each sub-directory is paired with +its own finished path. This is the parent :meth:`sftpretty.Connection.put_d` +appends a directory name to. Links are followed once per target, so a directory +pointing back at one of its own parents is mapped rather than chased. .. code-block:: python import sftpretty >>> directories = {} - >>> sftpretty.localtree(directories, '/home/user/downloads', '/tmp') + >>> sftpretty.helpers.localtree(directories, '/home/user/downloads', '/tmp') >>> directories - {'/home/user/downloads': [('/home/user/downloads/percona', '/tmp/downloads/percona'), - ('/home/user/downloads/wallstreet', '/tmp/downloads/wallstreet') - ] + {'/home/user/downloads': [('/home/user/downloads/percona', '/tmp/downloads'), + ('/home/user/downloads/wallstreet', '/tmp/downloads') + ], + '/home/user/downloads/wallstreet': [('/home/user/downloads/wallstreet/bets', + '/tmp/downloads/wallstreet') + ] } -:func:`sftpretty.st_mode_to_int` --------------------------------- -Converts an octal mode result back to an integer representation. The information -returned in SFTPAttribute object ``.stat(*fname*).st_mode`` contains extra -things you probably don't care about, in a form that has been converted from -octal to int so you won't recognize it at first. This function clips the extra -bits and hands you the file mode in a way you'll recognize. +:func:`sftpretty.helpers.retry` +------------------------------- +For the stubborn programmer in all of us. Calls sometimes fail for no good +reason and work on subsequent attempts. Name the exceptions worth another +attempt and wait ``delay`` seconds, multiplied by ``backoff``, then try again. +Specify a *type* to catch that whole family or an *instance* to catch exactly +one error while letting it siblings through. So ``IOError(errno.ECOMM)`` will +sit out a comms failure while a missing file still fails. Set ``silent`` to not +hear about it or ``logger`` to send the whole song and dance somewhere useful. +A value of 0 or None for ``tries`` returns your function undecorated. Its count +refers to total attempts rather than retries, so ``tries=3`` calls three times +at most. + +.. code-block:: python + + >>> from sftpretty.helpers import retry + + >>> @retry(TimeoutError, tries=3, delay=1, backoff=2) + ... def flaky(): + ... return connection.read() + + >>> flaky() + Retry (3/3): + connection reset + Retrying in 1 second(s)... + Retry (2/3): + connection reset + Retrying in 2 second(s)... + 'connected' + + +:func:`sftpretty.helpers.st_mode_to_int` +---------------------------------------- +Converts an octal mode result back to an integer representation. The +information returned in SFTPAttribute object ``.stat(*fname*).st_mode`` +contains extra things you probably don't care about, in a form that has been +converted from octal to int so you won't recognize it at first. This function +clips the extra bits and hands you the file mode in a way you'll recognize. .. code-block:: python >>> attr = sftp.stat('readme.txt') >>> attr.st_mode 33188 - >>> sftpretty.st_mode_to_int(attr.st_mode) + >>> sftpretty.helpers.st_mode_to_int(attr.st_mode) 644 diff --git a/pyproject.toml b/pyproject.toml index 6db4bed8..96eb0d50 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -48,7 +48,7 @@ keywords = [ name = 'sftpretty' readme = 'README.rst' requires-python = '>=3.6' -version = '1.2.3' +version = '1.2.4' [project.scripts] sftpretty = 'sftpretty:Connection' diff --git a/sftpretty/__init__.py b/sftpretty/__init__.py index 6fa7f5dc..532edaa5 100644 --- a/sftpretty/__init__.py +++ b/sftpretty/__init__.py @@ -145,11 +145,14 @@ def get_config(self, host): cval = self.ssh_config.lookup(host) return cval or {} - def get_hostkey(self, host): + def get_hostkey(self, host, key_type=None): '''Return the matching known hostkey to be used for verification or raise an SSHException. :param str host: The Hostname or IP of the remote machine. + :param str|None key_type: *Default: None* - Key type negotiated with + the remote such as ``ssh-ed25519``. When None the first hostkey + known for a host is returned. :returns: (obj) PKey - Public key(s) associated with host or None. @@ -159,6 +162,12 @@ def get_hostkey(self, host): # None | {key_type: private_key} if kval is None: raise SSHException(f'No hostkey for host [{host}] found.') + if key_type is not None: + hostkey = kval.get(key_type) + if hostkey is None: + raise SSHException(f'No [{key_type}] hostkey for host ' + f'[{host}] found.') + return hostkey # Return the public key from the dictionary return list(kval.values())[0] @@ -495,7 +504,7 @@ def _start_transport(self, host, port): if self._transport.is_active(): remote_hostkey = self._transport.get_remote_server_key() - remote_fingerprint = hash(remote_hostkey) + remote_fingerprint = hash(remote_hostkey.asbytes()) log.info((f'[{host}] Host Key: \n\t' f'Name: {remote_hostkey.get_name()}\n\t' f'Fingerprint: {remote_fingerprint}\n\t' @@ -507,8 +516,9 @@ def _start_transport(self, host, port): else: knownhost_name = host log.debug(f'Hostkey Name: {knownhost_name}') - user_hostkey = self._cnopts.get_hostkey(knownhost_name) - user_fingerprint = hash(user_hostkey) + user_hostkey = self._cnopts.get_hostkey( + knownhost_name, remote_hostkey.get_name()) + user_fingerprint = hash(user_hostkey.asbytes()) log.info(f'Known Fingerprint: {user_fingerprint}') if user_fingerprint != remote_fingerprint: raise HostKeysException((f'{host} key verification: ' @@ -521,6 +531,7 @@ def _start_transport(self, host, port): except (AttributeError, gaierror, UnicodeError): raise ConnectionException(host, port) except Exception as err: + self.close() raise err def get(self, remotefile, localpath=None, callback=None, diff --git a/sftpretty/helpers.py b/sftpretty/helpers.py index 719eada9..f82a52d4 100644 --- a/sftpretty/helpers.py +++ b/sftpretty/helpers.py @@ -1,6 +1,8 @@ +from errno import EBADF, ELOOP, ENOENT, ENOTDIR from functools import wraps from hashlib import new, sha3_512 from io import BytesIO, IOBase +from os import scandir from pathlib import Path, PureWindowsPath from re import sub from stat import S_IMODE @@ -8,6 +10,24 @@ def _callback(filename, bytes_so_far, bytes_total, logger=None): + '''log transfer progress as a percentage of the total + + :param str filename: + name of the file being transferred + :param int bytes_so_far: + bytes transferred so far + :param int bytes_total: + total bytes to transfer + :param logging.Logger logger: + logger instance to use. If None, print + + :returns: None + + :raises TypeError: + when bytes_so_far or bytes_total is not an integer + :raises ZeroDivisionError: + when bytes_total is zero + ''' message = (f'Transfer of File: [{filename}] @ ' f'{100.0 * bytes_so_far / bytes_total:.1f}% ' f'{bytes_so_far:d}:{bytes_total:d} bytes ') @@ -20,10 +40,13 @@ def _callback(filename, bytes_so_far, bytes_total, logger=None): def drivepath(filepath): '''Normalize a filepath to POSIX form, retaining any drive letter - :param str filename: + :param str filepath: path to file or string to process - :returns str: normalized POSIX path + :returns str: normalized POSIX path, empty input is passed through + + :raises TypeError: + when filepath is not a string ''' if filepath: if '\\' in filepath or PureWindowsPath(filepath).drive: @@ -53,17 +76,22 @@ def drivepath(filepath): def hash(filename, algorithm=sha3_512(), blocksize=65536): '''hash contents of a file, file like object or string - :param bytesIO,IObase,str filename: + :param BytesIO,IOBase,str filename: path to file, file object, or string to process :param hashlib.hash algorithm: hash object to use as digest algorithm :param int blocksize: size of chunk to read in avoiding memory exhaustion - :returns: hexdigest - - :raises: Exception + :returns str: hexdigest + :raises AttributeError: + when algorithm has no name attribute + :raises OSError: + when reading from a file object fails + :raises ValueError: + when algorithm.name is not a supported digest or filename is a + closed file object ''' buffer = new(algorithm.name) if isinstance(filename, str): @@ -73,6 +101,8 @@ def hash(filename, algorithm=sha3_512(), blocksize=65536): buffer.update(chunk) except OSError: buffer.update(bytes(filename.encode('utf-8'))) + elif isinstance(filename, bytes): + buffer.update(filename) elif isinstance(filename, BytesIO): for chunk in iter(lambda: filename.read1(blocksize), b''): buffer.update(chunk) @@ -84,46 +114,77 @@ def hash(filename, algorithm=sha3_512(), blocksize=65536): def localtree(container, localdir, remotedir, recurse=True): - '''recursively descend local directory mapping the tree to a - dictionary container. - - :param dict container: dictionary object to save directory tree - {localdir: - [(localdir/sub-directory, - remotedir/localdir/sub-directory)],} - {localdir: [(content path, remotedir/content path)],} + '''descend local directory mapping the tree to a dictionary container. + Subdirectories are paired with the remote directory they are created + in, not with their final path. Upstream function is responsible for + appending the name of the local directory it is handed to remotedir. + + :param dict container: + dictionary object to save directory tree + {localdir: [(localdir/sub-directory, remotedir/localdir)],} + {localdir: [(content path, remote parent of content path)],} :param str localdir: root of local directory to descend, use '.' to start at :attr:`.pwd` :param str remotedir: - root of remote directory to append localdir too - path - :param bool recurse: *Default: True*. To recurse or not to recurse - that is the question + root of remote directory localdir is created in + :param bool recurse: + *Default: True*. To recurse or not to recurse that is the + question :returns: None - :raises: Exception - + :raises AttributeError: + when localdir is not a string + :raises FileNotFoundError: + when localdir does not exist + :raises NotADirectoryError: + when localdir is not a directory + :raises PermissionError: + when a directory in the tree cannot be read ''' - try: - if localdir.startswith(':', 1) or localdir.startswith('\\'): - localdir = PureWindowsPath(localdir) - else: - localdir = Path(localdir).expanduser().absolute() - for localpath in Path(localdir).iterdir(): - if localpath.is_dir(): - local = localpath.as_posix() - remote = Path(remotedir).joinpath(localpath.relative_to( - localdir).as_posix()).as_posix() - if localdir.as_posix() in container.keys(): - container[localdir.as_posix()].append((local, remote)) - else: - container[localdir.as_posix()] = [(local, remote)] + if localdir.startswith(':', 1) or localdir.startswith('\\'): + localdir = Path(PureWindowsPath(localdir).as_posix()) + else: + localdir = Path(localdir).expanduser().absolute() + + branches = [(localdir.as_posix(), + Path(remotedir).joinpath(localdir.name).as_posix())] + seen = set() + + while branches: + branch = None + localroot, remotedir = branches.pop() + rootstat = None + + with scandir(localroot) as localpaths: + for localpath in localpaths: + try: + if not localpath.is_dir(): + continue + if localpath.is_symlink(): + if rootstat is None: + rootstat = Path(localroot).stat() + seen.add((rootstat.st_dev, rootstat.st_ino)) + symstat = localpath.stat() + softlink = (symstat.st_dev, symstat.st_ino) + if softlink in seen: + continue + seen.add(softlink) + except OSError as err: + if (err.errno in (EBADF, ELOOP, ENOENT, ENOTDIR) or + getattr(err, 'winerror', None) + in (21, 123, 1921)): + continue + raise + if branch is None: + branch = container.get(localroot) + if branch is None: + container[localroot] = branch = [] + local = f'{localroot}/{localpath.name}' + branch.append((local, remotedir)) if recurse: - localtree(container, local, remote, recurse=recurse) - except Exception as err: - raise err + branches.append((local, f'{remotedir}/{localpath.name}')) def retry(exceptions, tries=0, delay=3, backoff=2, silent=False, logger=None): @@ -134,20 +195,26 @@ def retry(exceptions, tries=0, delay=3, backoff=2, silent=False, logger=None): IOError or IOError(errno.ECOMM) or (IOError,) or (ValueError, IOError(errno.ECOMM) :param int tries: - number of times to try (not retry) before giving up. + number of times to try (not retry) before giving up :param int delay: - initial delay between retries in seconds. + initial delay between retries in seconds :param int backoff: - backoff multiplier. + backoff multiplier :param bool silent: if set then no logging will be attempted. - :param logging.logger logger: - logger instance to use. If None, print. - - :returns: wrapped function - - :raises: Exception - + :param logging.Logger logger: + logger instance to use. If None, print + + :returns function: + decorated function or the function unchanged when tries is None + or 0 + + :raises Exception: + whatever the decorated function raises, immediately when it is + not listed in exceptions, otherwise after tries are exhausted + :raises TypeError: + when exceptions holds anything that is not an exception type or + instance raised from the decorated call ''' try: len(exceptions) @@ -200,14 +267,17 @@ def _retry(*args, **kwargs): def st_mode_to_int(val): - '''SFTAttributes st_mode returns an stat type that shows more than what + '''SFTPAttributes st_mode returns an stat type that shows more than what can be set. Trim off those bits and convert to an int representation. - if you want an object that was `chmod 711` to return a value of 711, use - this function + If you want an object that was `chmod 711` to return a value of 711, use + this function. - :param int val: the value of an st_mode attr returned by SFTPAttributes + :param int val: + the value of an st_mode attr returned by SFTPAttributes :returns int: integer representation of octal mode + :raises TypeError: + when val is not an integer ''' return int(str(oct(S_IMODE(val)))[-3:]) diff --git a/tests/common.py b/tests/common.py index 8508863a..6e348b6d 100644 --- a/tests/common.py +++ b/tests/common.py @@ -2,7 +2,6 @@ import pytest -from contextlib import contextmanager from os import environ from pathlib import Path from sftpretty import CnOpts diff --git a/tests/conftest.py b/tests/conftest.py index 3545d759..db678439 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -3,9 +3,17 @@ import pytest from common import LOCAL, remote_rmdir, STARS8192, USER_HOME +from cryptography.hazmat.primitives.asymmetric.ed25519 import ( + Ed25519PrivateKey) +from cryptography.hazmat.primitives.serialization import (Encoding, + NoEncryption, + PrivateFormat) +from io import StringIO from os import close +from paramiko import Ed25519Key from paramiko.hostkeys import HostKeys from pathlib import Path +from pytest_sftpserver.sftp.server import SFTPRequestHandler from sftpretty import CnOpts, Connection from tempfile import mkstemp from uuid import uuid4 @@ -23,18 +31,20 @@ def lsftp(request): @pytest.fixture(autouse=True, scope='module') -def knownhosts(sftpserver, key_type='ssh-ed25519'): +def knownhosts(sftpserver): '''setup host key for test server in local knownhosts''' if sftpserver.port != 22: host = f'[{sftpserver.host}]:{sftpserver.port}' else: host = sftpserver.host host_hashed = HostKeys().hash_host(host) - hostkey = \ - 'AAAAC3NzaC1lZDI1NTE5AAAAIB0g3SG/bbyysJ7f0kqdoWMXhHxxFR7aLJYNIHO/MtsD' + private = Ed25519PrivateKey.generate().private_bytes( + Encoding.PEM, PrivateFormat.OpenSSH, NoEncryption()) + hostkey = Ed25519Key(file_obj=StringIO(private.decode())) + SFTPRequestHandler.host_key = hostkey hostkeys = f'''\ - {host} {key_type} {hostkey} - {host_hashed} {key_type} {hostkey}''' + {host} {hostkey.get_name()} {hostkey.get_base64()} + {host_hashed} {hostkey.get_name()} {hostkey.get_base64()}''' knownhosts = Path('sftpserver.pub') knownhosts.write_bytes(bytes(hostkeys, 'utf-8')) diff --git a/tests/test_helpers.py b/tests/test_helpers.py index b1cb5304..437f07d2 100644 --- a/tests/test_helpers.py +++ b/tests/test_helpers.py @@ -4,10 +4,42 @@ from hashlib import md5, new, sha1, sha256, sha3_512 from io import BytesIO -from sftpretty.helpers import drivepath, hash +from logging import getLogger, INFO +from sftpretty.helpers import _callback, drivepath, hash, st_mode_to_int +from types import SimpleNamespace -@pytest.mark.parametrize('path,expected', ( +@pytest.mark.parametrize('current, total, percent', ( + (1024, 1024, '100.0%'), + (512, 1024, '50.0%'), + (1, 3, '33.3%'), + (0, 1024, '0.0%')), + ids=('complete', 'half', 'rounded', 'start')) +def test_callback(current, total, percent, capsys): + '''test progress prints as a percentage of the total''' + _callback('eels.txt', current, total) + printed = capsys.readouterr().out + + assert 'eels.txt' in printed + assert percent in printed + assert f'{current}:{total}' in printed + + +def test_callback_logger(caplog): + '''test progress is logged when a logger is provided''' + with caplog.at_level(INFO): + _callback('eels.txt', 512, 1024, logger=getLogger('sftpretty')) + + assert '50.0%' in caplog.text + + +def test_callback_zero(): + '''test a zero byte total raises''' + with pytest.raises(ZeroDivisionError): + _callback('eels.txt', 0, 0) + + +@pytest.mark.parametrize('path, expected', ( # drive qualified ('C:\tmp\test.txt', '/C:/tmp/test.txt'), ('C:\\tmp\test.txt', '/C:/tmp/test.txt'), @@ -25,7 +57,7 @@ ('/C:', '/C:/'), ('/C:/', '/C:/'), # leading backslash is UNC, typed pairs arrive collapsed ('\\\\server\\share\\file.txt', '//server/share/file.txt'), - ('\\server\share\file.txt', '//server/share\file.txt'), # noqa: W605 + ('\\server\\share\file.txt', '//server/share\file.txt'), ('\\tmp\test.txt', '//tmp/test.txt'), ('//tmp/test.txt', '//tmp/test.txt'), ('//server/share//dbl/f.txt', '//server/share/dbl/f.txt'), @@ -50,6 +82,14 @@ def test_drivepath(path, expected): assert drivepath(expected) == expected +@pytest.mark.parametrize('path', (1024, b'C:\\tmp\\test.txt'), + ids=('integer', 'bytes')) +def test_drivepath_unsupported(path): + '''test a non string argument raises''' + with pytest.raises(TypeError): + drivepath(path) + + @pytest.mark.parametrize('algorithm', (md5(), sha1(), sha256(), sha3_512()), ids=('md5', 'sha1', 'sha256', 'sha3_512')) def test_hash_algorithm(algorithm, tempfile_containing): @@ -62,6 +102,18 @@ def test_hash_algorithm(algorithm, tempfile_containing): assert hash(content, algorithm=algorithm) == expected +def test_hash_algorithm_nameless(): + '''test an algorithm without a name attribute raises''' + with pytest.raises(AttributeError): + hash('some content', algorithm=object()) + + +def test_hash_algorithm_unknown(): + '''test an unsupported digest name raises''' + with pytest.raises(ValueError): + hash('some content', algorithm=SimpleNamespace(name='nonesuch')) + + @pytest.mark.parametrize('blocksize', (1, 7, 65536), ids=('byte', 'partial', 'default')) def test_hash_blocksize(blocksize, tempfile_containing): @@ -74,6 +126,15 @@ def test_hash_blocksize(blocksize, tempfile_containing): assert hash(BytesIO(content.encode()), blocksize=blocksize) == expected +def test_hash_closed(tempfile_containing): + '''test a closed file object raises''' + filestream = open(tempfile_containing(contents='some content'), 'rb') + filestream.close() + + with pytest.raises(ValueError): + hash(filestream) + + @pytest.mark.parametrize('left, right', ( ('CASE', 'case'), ('some content', 'some different content'), @@ -133,3 +194,40 @@ def test_hash_repeatable(tempfile_containing): def test_hash_unreadable(unreadable): '''test a string that cannot be opened is digested as a string''' assert hash(unreadable) == sha3_512(unreadable.encode()).hexdigest() + + +@pytest.mark.xfail(reason='unsupported types digest as an empty buffer') +def test_hash_unsupported(): + '''test input that is neither a path, string nor file object raises''' + with pytest.raises(TypeError): + hash(1024) + + +@pytest.mark.parametrize('mode,expected', ( + (0o100400, 400), (0o100644, 644), (0o100711, 711), + (0o40755, 755), (0o40777, 777)), + ids=('400', '644', '711', '755', '777')) +def test_st_mode_to_int(mode, expected): + '''test file type bits are trimmed from the mode''' + assert st_mode_to_int(mode) == expected + + +@pytest.mark.parametrize('mode,expected', ((0o41777, 777), (0o104755, 755)), + ids=('sticky', 'setuid')) +def test_st_mode_to_int_special(mode, expected): + '''test set user ID, set group ID and sticky bits are dropped''' + assert st_mode_to_int(mode) == expected + + +@pytest.mark.parametrize('val', ('0755', None, 7.55), + ids=('string', 'none', 'float')) +def test_st_mode_to_int_unsupported(val): + '''test a non integer mode raises''' + with pytest.raises(TypeError): + st_mode_to_int(val) + + +@pytest.mark.xfail(raises=ValueError, reason='oct(0) renders as 0o0') +def test_st_mode_to_int_zero(): + '''test a mode carrying no permission bits converts to zero''' + assert st_mode_to_int(0o100000) == 0 diff --git a/tests/test_localtree.py b/tests/test_localtree.py index 18a43c29..3bfe1382 100644 --- a/tests/test_localtree.py +++ b/tests/test_localtree.py @@ -1,66 +1,247 @@ -'''test sftpretty.localtree''' +'''test sftpretty.helpers.localtree''' -from common import conn, rmdir, VFS +import pytest + +from blddirs import build_dir_struct from pathlib import Path -from sftpretty import Connection, localtree -from tempfile import mkdtemp +from sftpretty.helpers import localtree +from sys import getrecursionlimit -def test_localtree(sftpserver): +def test_localtree(tmp_path): '''test the localtree function, with recurse''' - with sftpserver.serve_content(VFS): - with Connection(**conn(sftpserver)) as sftp: - localpath = Path(mkdtemp()).as_posix() - sftp.get_r('.', localpath) - - cwd = sftp.pwd - tree = {} - - localtree(tree, localpath, cwd) - - local = { - f'{localpath}': [ - (f'{localpath}/pub', cwd + '/pub') - ], - f'{localpath}/pub': [ - (f'{localpath}/pub/foo1', cwd + '/pub/foo1'), - (f'{localpath}/pub/foo2', cwd + '/pub/foo2') - ], - f'{localpath}/pub/foo2': [ - (f'{localpath}/pub/foo2/bar1', cwd + '/pub/foo2/bar1') - ] - } - - for branch in sorted(tree.keys()): - assert set(local[branch]) == set(tree[branch]) - del tree[branch] - assert tree == {} - - rmdir(localpath) - - -def test_localtree_no_recurse(sftpserver): - '''test the localtree function, without recursing''' - with sftpserver.serve_content(VFS): - with Connection(**conn(sftpserver)) as sftp: - localpath = Path(mkdtemp()).as_posix() - sftp.chdir('pub/foo2') - sftp.get_r('.', localpath) - - cwd = sftp.pwd - tree = {} - - localtree(tree, localpath, cwd, recurse=False) - - local = { - f'{localpath}': [ - (f'{localpath}/bar1', cwd + '/bar1') - ] - } - - for branch in sorted(tree.keys()): - assert set(local[branch]) == set(tree[branch]) - del tree[branch] - assert tree == {} - - rmdir(localpath) + build_dir_struct(tmp_path.as_posix()) + container = {} + local = tmp_path.joinpath('pub').as_posix() + + localtree(container, local, '/remote') + + expected = { + local: [(f'{local}/foo1', '/remote/pub'), + (f'{local}/foo2', '/remote/pub')], + f'{local}/foo2': [(f'{local}/foo2/bar1', '/remote/pub/foo2')] + } + + assert container.keys() == expected.keys() + for branch in expected: + assert set(container[branch]) == set(expected[branch]) + + +def test_localtree_deep(tmp_path): + '''test a tree deeper than the recursion limit is mapped''' + container = {} + levels = getrecursionlimit() + 100 + local = tmp_path.joinpath('deep') + local.mkdir() + + branch = local + try: + for _ in range(levels): + try: + branch.joinpath('d').mkdir() + except OSError: + pytest.skip('platform path limit reached') + branch = branch.joinpath('d') + + localtree(container, local.as_posix(), '/remote') + + assert len(container) == levels + finally: + while branch != local: + branch.rmdir() + branch = branch.parent + + +def test_localtree_hidden(tmp_path): + '''test dot directories are mapped''' + build_dir_struct(tmp_path.as_posix()) + container = {} + local = tmp_path.joinpath('pub').as_posix() + tmp_path.joinpath('pub', '.hidden').mkdir() + + localtree(container, local, '/remote', recurse=False) + + assert (f'{local}/.hidden', '/remote/pub') in container[local] + + +def test_localtree_leaf(tmp_path): + '''test a directory holding no sub-directories creates no key''' + build_dir_struct(tmp_path.as_posix()) + container = {} + + localtree(container, tmp_path.joinpath('pub', 'foo1').as_posix(), + '/remote') + + assert container == {} + + +def test_localtree_missing(tmp_path): + '''test a local directory that does not exist raises''' + with pytest.raises(FileNotFoundError): + localtree({}, tmp_path.joinpath('nonesuch').as_posix(), '/remote') + + +def test_localtree_no_recurse(tmp_path): + '''test only the first level is mapped without recursing''' + build_dir_struct(tmp_path.as_posix()) + container = {} + local = tmp_path.joinpath('pub').as_posix() + + localtree(container, local, '/remote', recurse=False) + + assert container.keys() == {local} + assert set(container[local]) == {(f'{local}/foo1', '/remote/pub'), + (f'{local}/foo2', '/remote/pub')} + + +def test_localtree_notadirectory(tempfile_containing): + '''test a file in place of a local directory raises''' + with pytest.raises(NotADirectoryError): + localtree({}, tempfile_containing(contents='eels'), '/remote') + + +def test_localtree_parents_first(tmp_path): + '''test a branch is always mapped before any of its children''' + build_dir_struct(tmp_path.as_posix()) + container = {} + tmp_path.joinpath('pub', 'foo2', 'bar1', 'baz1').mkdir() + + localtree(container, tmp_path.joinpath('pub').as_posix(), '/remote') + + branches = list(container) + for index, branch in enumerate(branches): + parents = [parent for parent in branches + if branch.startswith(f'{parent}/')] + + assert all(branches.index(parent) < index for parent in parents) + + +def test_localtree_put_d_contract(tmp_path): + '''test the directory name appended to its remote is the final path''' + build_dir_struct(tmp_path.as_posix()) + container = {} + + localtree(container, tmp_path.joinpath('pub').as_posix(), '/remote') + + for branch in container.values(): + for path, remote in branch: + relative = Path(path).relative_to(tmp_path).as_posix() + + assert f'{remote}/{Path(path).name}' == f'/remote/{relative}' + + +def test_localtree_relative(tmp_path, monkeypatch): + '''test a relative localdir descends the working directory''' + build_dir_struct(tmp_path.as_posix()) + container = {} + monkeypatch.chdir(tmp_path.joinpath('pub')) + + localtree(container, '.', '/remote', recurse=False) + + assert container.keys() == {Path.cwd().as_posix()} + + +@pytest.mark.parametrize('remotedir, expected', ( + ('/', '/pub'), + ('/remote', '/remote/pub'), + ('/remote/', '/remote/pub')), + ids=('root', 'plain', 'trailing')) +def test_localtree_remotedir(remotedir, expected, tmp_path): + '''test remote roots join without doubling the separator''' + build_dir_struct(tmp_path.as_posix()) + container = {} + local = tmp_path.joinpath('pub').as_posix() + + localtree(container, local, remotedir, recurse=False) + + assert {remote for _, remote in container[local]} == {expected} + + +def test_localtree_seeded(tmp_path): + '''test a container seeded by put_r is extended, not replaced''' + build_dir_struct(tmp_path.as_posix()) + local = tmp_path.joinpath('pub').as_posix() + container = {local: [(local, '/remote')]} + + localtree(container, local, '/remote') + + assert (local, '/remote') in container[local] + assert len(container[local]) == 3 + + +def test_localtree_symlink(tmp_path): + '''test a link pointing outside the tree is followed''' + build_dir_struct(tmp_path.as_posix()) + container = {} + local = tmp_path.joinpath('pub').as_posix() + tmp_path.joinpath('outside', 'deep').mkdir(parents=True) + tmp_path.joinpath('pub', 'link').symlink_to( + tmp_path.joinpath('outside'), target_is_directory=True) + + localtree(container, local, '/remote') + + assert (f'{local}/link', '/remote/pub') in container[local] + assert (f'{local}/link/deep', + '/remote/pub/link') in container[f'{local}/link'] + + +def test_localtree_symlink_dangling(tmp_path): + '''test a link with no target is skipped''' + build_dir_struct(tmp_path.as_posix()) + container = {} + local = tmp_path.joinpath('pub').as_posix() + gone = tmp_path.joinpath('pub', 'nonesuch') + gone.mkdir() + tmp_path.joinpath('pub', 'gone').symlink_to(gone, + target_is_directory=True) + gone.rmdir() + + localtree(container, local, '/remote', recurse=False) + + assert f'{local}/gone' not in [path for path, _ in container[local]] + + +def test_localtree_symlink_duplicate(tmp_path): + '''test the same target is only followed once''' + build_dir_struct(tmp_path.as_posix()) + container = {} + local = tmp_path.joinpath('pub').as_posix() + for link in ('one', 'two'): + tmp_path.joinpath('pub', link).symlink_to( + tmp_path.joinpath('pub', 'foo1'), target_is_directory=True) + + localtree(container, local, '/remote') + + assert len([path for path, _ in container[local] + if path.endswith(('/one', '/two'))]) == 1 + + +def test_localtree_symlink_loop(tmp_path): + '''test a link to its own parent is never descended''' + build_dir_struct(tmp_path.as_posix()) + container = {} + local = tmp_path.joinpath('pub').as_posix() + tmp_path.joinpath('pub', 'cycle').symlink_to( + tmp_path.joinpath('pub'), target_is_directory=True) + + localtree(container, local, '/remote') + + assert f'{local}/cycle' not in [path for path, _ in container[local]] + assert container.keys() == {local, f'{local}/foo2'} + + +def test_localtree_trailing_slash(tmp_path): + '''test a trailing separator on localdir is absorbed''' + build_dir_struct(tmp_path.as_posix()) + container = {} + local = tmp_path.joinpath('pub').as_posix() + + localtree(container, f'{local}/', '/remote', recurse=False) + + assert container.keys() == {local} + + +def test_localtree_unsupported(tmp_path): + '''test a non string localdir raises''' + with pytest.raises(AttributeError): + localtree({}, tmp_path, '/remote') diff --git a/tests/test_retry.py b/tests/test_retry.py index 82a20060..60cee973 100644 --- a/tests/test_retry.py +++ b/tests/test_retry.py @@ -1,14 +1,16 @@ +'''test sftpretty.helpers.retry''' + import pytest from logging import DEBUG, getLogger, StreamHandler from sftpretty.helpers import retry -class RetryableError(Exception): +class AnotherRetryableError(Exception): pass -class AnotherRetryableError(Exception): +class RetryableError(Exception): pass @@ -16,25 +18,35 @@ class UnexpectedError(Exception): pass -def test_no_retry_required(): - counter = 0 +def test_disabled_returns_undecorated(): - @retry(RetryableError, tries=4, delay=0.1) def succeeds(): + return 'success' + + assert retry(RetryableError, silent=True)(succeeds) is succeeds + assert retry(RetryableError, tries=0, silent=True)(succeeds) is succeeds + assert retry(RetryableError, tries=None, silent=True)(succeeds) is succeeds + + +def test_exception_instance_ignores_other_args(): + counter = 0 + + @retry(RetryableError('failed'), tries=4, delay=0.1) + def raises_other_args(): nonlocal counter counter += 1 - return 'success' + raise RetryableError('a different message') - r = succeeds() + with pytest.raises(RetryableError, match='a different message'): + raises_other_args() - assert r == 'success' assert counter == 1 -def test_retries_once(): +def test_exception_instance_matches_args(): counter = 0 - @retry(RetryableError, tries=4, delay=0.1) + @retry(RetryableError('failed'), tries=4, delay=0.1) def fails_once(): nonlocal counter counter += 1 @@ -49,6 +61,16 @@ def fails_once(): assert counter == 2 +def test_invalid_exception_type_raises(): + + @retry('failed', tries=4, delay=0.1) + def raise_retryable_error(): + raise RetryableError('failed') + + with pytest.raises(TypeError): + raise_retryable_error() + + def test_limit_is_reached(): counter = 0 @@ -84,6 +106,39 @@ def raise_multiple_exceptions(): assert counter == 3 +def test_no_retry_required(): + counter = 0 + + @retry(RetryableError, tries=4, delay=0.1) + def succeeds(): + nonlocal counter + counter += 1 + return 'success' + + r = succeeds() + + assert r == 'success' + assert counter == 1 + + +def test_retries_once(): + counter = 0 + + @retry(RetryableError, tries=4, delay=0.1) + def fails_once(): + nonlocal counter + counter += 1 + if counter < 2: + raise RetryableError('failed') + else: + return 'success' + + r = fails_once() + + assert r == 'success' + assert counter == 2 + + def test_unexpected_exception_does_not_retry(): @retry(RetryableError, tries=4, delay=0.1) @@ -94,7 +149,6 @@ def raise_unexpected_error(): raise_unexpected_error() -@pytest.fixture(autouse=True) def test_using_a_logger(caplog): _caplog = caplog counter = 0