← 返回 Skill 说明

v2.1.0 / Python 工具

工具源码

file_guard.py

查看原始文件 ↗
#!/usr/bin/env python3
"""Read-only directory inventory/compare. Python 3.10+, POSIX, standard library.

Root directories and their ancestors must be real directories. Child symlinks
are recorded without traversal. Metadata is checked while reading and in a
second complete pass; this detects observed instability, NOT an atomic snapshot.
Outputs contain private paths and must be kept private. No cleanup capability.
"""
import argparse
import errno
import hashlib
import json
import os
import re
import stat
import sys

VERSION = 1
LABEL = re.compile(r"[A-Za-z][A-Za-z0-9_-]*\Z")
SHA = re.compile(r"[0-9a-f]{64}\Z")
DIR_FLAGS = os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW


class Invalid(Exception):
    """Only a fixed, non-sensitive diagnostic code is exposed to the CLI."""
    def __init__(self, code='INVALID_ARGUMENTS'):
        self.code = code
        super().__init__(code)


def require(ok, code='INVALID_ARGUMENTS'):
    if not ok:
        raise Invalid(code)


def absolute(path):
    require(isinstance(path, str) and path.startswith('/') and '\0' not in path)
    require(path == os.path.normpath(path) and not path.startswith('//'))
    return path


def open_directory(path):
    """Open each component relative to a held descriptor; never follow links."""
    absolute(path)
    fd = os.open('/', DIR_FLAGS)
    try:
        for part in path.split('/')[1:]:
            if part:
                child = os.open(part, DIR_FLAGS, dir_fd=fd)
                os.close(fd)
                fd = child
        return fd
    except OSError as error:
        os.close(fd)
        if error.errno in (errno.ELOOP, errno.ENOTDIR):
            raise Invalid('PATH_NOT_REAL_DIRECTORY') from None
        raise
    except BaseException:
        os.close(fd)
        raise


def fingerprint(s):
    # atime is deliberately excluded: reading may update it.
    return (s.st_dev, s.st_ino, s.st_mode, s.st_size, s.st_mtime_ns,
            s.st_ctime_ns, s.st_nlink, s.st_uid, s.st_gid)


def validate_roots(roots):
    require(isinstance(roots, dict) and bool(roots))
    paths = []
    for label, path in roots.items():
        require(isinstance(label, str) and LABEL.fullmatch(label))
        absolute(path)
        for other in paths:
            require(os.path.commonpath([path, other]) not in (path, other))
        paths.append(path)


def scan(roots, hash_files):
    entries, identities = {}, {}

    def visit(parent, name, key, root=False):
        before = os.fstat(parent) if root else os.stat(name, dir_fd=parent, follow_symlinks=False)
        identity = fingerprint(before)
        identities[key] = identity
        mode = stat.S_IMODE(before.st_mode)
        if stat.S_ISDIR(before.st_mode):
            fd = os.dup(parent) if root else os.open(name, DIR_FLAGS, dir_fd=parent)
            try:
                require(fingerprint(os.fstat(fd)) == identity, 'SOURCE_CHANGED')
                entries[key] = {'type': 'directory', 'mode': mode}
                names = sorted(os.listdir(fd))
                for child in names:
                    visit(fd, child, key + ('' if key.endswith('/') else '/') + child)
                require(sorted(os.listdir(fd)) == names, 'SOURCE_CHANGED')
                require(fingerprint(os.fstat(fd)) == identity, 'SOURCE_CHANGED')
            finally:
                os.close(fd)
        elif stat.S_ISREG(before.st_mode):
            # NONBLOCK prevents a replaced FIFO from blocking before fstat.
            fd = os.open(name, os.O_RDONLY | os.O_NOFOLLOW | os.O_NONBLOCK, dir_fd=parent)
            try:
                require(fingerprint(os.fstat(fd)) == identity, 'SOURCE_CHANGED')
                digest = hashlib.sha256()
                if hash_files:
                    while True:
                        chunk = os.read(fd, 1024 * 1024)
                        if not chunk:
                            break
                        digest.update(chunk)
                require(fingerprint(os.fstat(fd)) == identity, 'SOURCE_CHANGED')
                entries[key] = {'type': 'file', 'mode': mode, 'size': before.st_size,
                                'sha256': digest.hexdigest()}
            finally:
                os.close(fd)
        elif stat.S_ISLNK(before.st_mode):
            entries[key] = {'type': 'symlink', 'mode': mode,
                            'target': os.readlink(name, dir_fd=parent)}
        else:
            raise Invalid('UNSUPPORTED_ENTRY')
        after = os.fstat(parent) if root else os.stat(name, dir_fd=parent, follow_symlinks=False)
        require(fingerprint(after) == identity, 'SOURCE_CHANGED')

    for label, path in sorted(roots.items()):
        fd = open_directory(path)
        try:
            visit(fd, '', label + '/', root=True)
        except OSError as error:
            code = ('SOURCE_CHANGED' if error.errno in
                    (errno.ENOENT, errno.ELOOP, errno.ENOTDIR) else 'SOURCE_UNREADABLE')
            raise Invalid(code) from None
        finally:
            os.close(fd)
    return entries, identities


def outside(path, roots):
    absolute(path)
    for root in roots.values():
        require(os.path.commonpath([path, root]) != root, 'OUTPUT_INSIDE_SOURCE')


def _write_private(path, data, roots):
    outside(path, roots)
    parent, name = os.path.split(path)
    require(bool(name))
    fd = open_directory(parent)
    try:
        out = os.open(name, os.O_WRONLY | os.O_CREAT | os.O_EXCL | os.O_NOFOLLOW,
                      0o600, dir_fd=fd)
        # Keep a partial output on I/O failure; never overwrite or delete it.
        with os.fdopen(out, 'w', encoding='utf-8') as stream:
            os.fchmod(stream.fileno(), 0o600)
            json.dump(data, stream, ensure_ascii=True, indent=2, sort_keys=True)
            stream.write('\n')
            stream.flush()
            os.fsync(stream.fileno())
    finally:
        os.close(fd)


def write_private(path, data, roots):
    try:
        _write_private(path, data, roots)
    except FileExistsError:
        raise Invalid('OUTPUT_EXISTS') from None
    except OSError:
        raise Invalid('OUTPUT_UNAVAILABLE') from None
    except Invalid as error:
        if error.code == 'OUTPUT_INSIDE_SOURCE':
            raise
        raise Invalid('OUTPUT_UNAVAILABLE') from None


def unique_object(pairs):
    result = {}
    for key, value in pairs:
        require(key not in result)
        result[key] = value
    return result


def read_json(path):
    absolute(path)
    parent, name = os.path.split(path)
    fd = open_directory(parent)
    try:
        source = os.open(name, os.O_RDONLY | os.O_NOFOLLOW | os.O_NONBLOCK, dir_fd=fd)
        with os.fdopen(source, 'r', encoding='utf-8') as stream:
            initial = os.fstat(stream.fileno())
            require(stat.S_ISREG(initial.st_mode))
            value = json.load(stream, object_pairs_hook=unique_object)
            require(fingerprint(initial) == fingerprint(os.fstat(stream.fileno())))
            require(fingerprint(initial) == fingerprint(os.stat(name, dir_fd=fd, follow_symlinks=False)))
            return value
    finally:
        os.close(fd)


def manifest(value):
    require(isinstance(value, dict) and set(value) == {'version', 'roots', 'entries'})
    require(type(value['version']) is int and value['version'] == VERSION)
    validate_roots(value['roots'])
    entries = value['entries']
    require(isinstance(entries, dict))
    for key, entry in entries.items():
        require(isinstance(key, str) and '/' in key and '\0' not in key)
        label, relative = key.split('/', 1)
        require(label in value['roots'])
        require(not relative or all(p not in ('', '.', '..') for p in relative.split('/')))
        require(isinstance(entry, dict) and type(entry.get('mode')) is int and 0 <= entry['mode'] <= 0o7777)
        kind = entry.get('type')
        fields = {'type', 'mode'}
        if kind == 'file':
            fields |= {'size', 'sha256'}
            require(type(entry.get('size')) is int and entry['size'] >= 0)
            require(isinstance(entry.get('sha256'), str) and SHA.fullmatch(entry['sha256']))
        elif kind == 'symlink':
            fields |= {'target'}
            require(isinstance(entry.get('target'), str) and bool(entry['target']) and '\0' not in entry['target'])
        else:
            require(kind == 'directory')
        require(set(entry) == fields)
        if relative:
            parent = key.rsplit('/', 1)[0] if '/' in relative else label + '/'
            require(parent in entries and isinstance(entries[parent], dict) and entries[parent].get('type') == 'directory')
    for label in value['roots']:
        require(entries.get(label + '/', {}).get('type') == 'directory')
    return value


def run(args):
    if args.command == 'snapshot':
        roots = {}
        for item in args.root:
            require('=' in item)
            label, path = item.split('=', 1)
            require(label not in roots)
            roots[label] = path
        validate_roots(roots)
        outside(args.out, roots)
        try:
            entries, first = scan(roots, True)
        except OSError:
            raise Invalid('SOURCE_UNREADABLE') from None
        try:
            _, second = scan(roots, False)
        except (Invalid, OSError):
            # A previously readable tree can no longer be verified unchanged.
            raise Invalid('SOURCE_CHANGED') from None
        require(first == second, 'SOURCE_CHANGED')
        write_private(args.out, {'version': VERSION, 'roots': roots, 'entries': entries}, roots)
        print('SNAPSHOT entries=' + str(len(entries)))
        return 0
    try:
        before, after = manifest(read_json(args.before)), manifest(read_json(args.after))
        require(before['roots'] == after['roots'])
    except (Invalid, OSError, ValueError, TypeError, RecursionError):
        raise Invalid('INVALID_MANIFEST') from None
    roots = before['roots']
    outside(args.out, roots)
    old, new = before['entries'], after['entries']
    missing = sorted(old.keys() - new.keys())
    added = sorted(new.keys() - old.keys())
    changed = sorted(k for k in old.keys() & new.keys() if old[k] != new[k])
    try:
        allowed = read_json(args.allow_removed) if args.allow_removed else []
        require(isinstance(allowed, list) and all(isinstance(k, str) for k in allowed))
        require(len(allowed) == len(set(allowed)) and set(allowed) <= set(missing))
    except (Invalid, OSError, ValueError, TypeError, RecursionError):
        raise Invalid('INVALID_REMOVAL_LIST') from None
    unexpected = sorted(set(missing) - set(allowed))
    passed = not (unexpected or added or changed)
    write_private(args.out, {'version': VERSION, 'status': 'PASS' if passed else 'DIFF',
                            'roots': roots, 'missing': missing, 'added': added,
                            'changed': changed, 'allowed_removed': sorted(allowed),
                            'unexpected_missing': unexpected}, roots)
    print(('PASS' if passed else 'DIFF') + ' changed=' + str(len(changed)) +
          ' missing=' + str(len(missing)) + ' added=' + str(len(added)) +
          ' allowed_removed=' + str(len(allowed)))
    return 0 if passed else 1


class QuietParser(argparse.ArgumentParser):
    def error(self, message):
        self.exit(2, 'ERROR: INVALID_ARGUMENTS; use --help.\n')


def main():
    parser = QuietParser(description=__doc__)
    sub = parser.add_subparsers(dest='command', required=True)
    snapshot = sub.add_parser('snapshot')
    snapshot.add_argument('--root', action='append', required=True, metavar='LABEL=/absolute/path')
    snapshot.add_argument('--out', required=True)
    compare = sub.add_parser('compare')
    compare.add_argument('--before', required=True)
    compare.add_argument('--after', required=True)
    compare.add_argument('--out', required=True)
    compare.add_argument('--allow-removed', help='Human-issued exact-key JSON list; not authorization by itself')
    args = parser.parse_args()
    try:
        return run(args)
    except Invalid as error:
        print('ERROR: ' + error.code + '; no verified result.', file=sys.stderr)
        return 2
    except (OSError, ValueError, TypeError, RecursionError):
        print('ERROR: INPUT_UNAVAILABLE; no verified result.', file=sys.stderr)
        return 2


if __name__ == '__main__':
    sys.exit(main())