← 返回 Skill 说明

v2.1.0 / Python 工具

工具源码

json_candidate.py

查看原始文件 ↗
#!/usr/bin/env python3
"""Bounded object-member removal; Python 3.10+ POSIX, standard library.

Default preview. --candidate requires --expected-sha256. --targets is a private
0600 JSON file under an owned 0700 directory, containing exact paths, e.g.
[["account", "oldToken"], ["oldPreference"]]. Only object keys are traversed.
No source replacement, locking, atomic snapshot, or application-state checks.
Output must use a different parent directory. The caller must choose a workspace
outside all application-managed roots; the helper cannot discover those roots.
Raw surviving tokens and order are preserved; only removal spans are deleted.
Whitespace inside those spans is removed as needed to keep JSON valid.
"""
import argparse
from dataclasses import dataclass
import hashlib
import json
import os
import re
import stat
import sys

MAX_BYTES = 8 * 1024 * 1024
MAX_DEPTH = 64
MAX_NODES = 100000
MAX_TARGETS = 256
DIR_FLAGS = os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW
FILE_FLAGS = os.O_RDONLY | os.O_NOFOLLOW | os.O_NONBLOCK
NUMBER = re.compile(r'-?(?:0|[1-9][0-9]*)(?:\.[0-9]+)?(?:[eE][+-]?[0-9]+)?')
SHA = re.compile(r'[0-9a-f]{64}\Z')


class Refusal(Exception):
    def __init__(self, code):
        self.code = code
        super().__init__(code)


def require(condition, code):
    if not condition:
        raise Refusal(code)


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


def directory(path):
    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 BaseException:
        os.close(fd)
        raise


def identity(info):
    return (info.st_dev, info.st_ino, info.st_mode, info.st_size,
            info.st_mtime_ns, info.st_ctime_ns, info.st_uid, info.st_gid,
            info.st_nlink)


def private(info, mode):
    return info.st_uid == os.getuid() and stat.S_IMODE(info.st_mode) == mode


def read_file(path, private_input=False):
    """Check observed identity and bytes on two reads; do not follow symlinks."""
    absolute(path)
    parent, name = os.path.split(path)
    require(bool(name), 'INVALID_PATH')
    folder = directory(parent)
    try:
        if private_input:
            require(private(os.fstat(folder), 0o700), 'TARGETS_NOT_PRIVATE')
        before = os.stat(name, dir_fd=folder, follow_symlinks=False)
        require(stat.S_ISREG(before.st_mode), 'SOURCE_NOT_REGULAR')
        if private_input:
            require(private(before, 0o600), 'TARGETS_NOT_PRIVATE')
        fd = os.open(name, FILE_FLAGS, dir_fd=folder)
        try:
            require(identity(os.fstat(fd)) == identity(before), 'SOURCE_CHANGED')
            require(before.st_size <= MAX_BYTES, 'LIMIT_EXCEEDED')
            chunks, total = [], 0
            while True:
                chunk = os.read(fd, min(65536, MAX_BYTES + 1 - total))
                if not chunk:
                    break
                total += len(chunk)
                require(total <= MAX_BYTES, 'LIMIT_EXCEEDED')
                chunks.append(chunk)
            data = b''.join(chunks)
            os.lseek(fd, 0, os.SEEK_SET)
            offset = 0
            while True:
                chunk = os.read(fd, 65536)
                if not chunk:
                    break
                require(data[offset:offset + len(chunk)] == chunk, 'SOURCE_CHANGED')
                offset += len(chunk)
            require(offset == len(data), 'SOURCE_CHANGED')
            require(identity(os.fstat(fd)) == identity(before), 'SOURCE_CHANGED')
            require(identity(os.stat(name, dir_fd=folder, follow_symlinks=False))
                    == identity(before), 'SOURCE_CHANGED')
            return data, identity(before)
        finally:
            os.close(fd)
    finally:
        os.close(folder)


@dataclass
class Node:
    kind: str
    start: int
    end: int
    children: object = None
    commas: object = None


class Parser:
    """Strict JSON parser recording spans without converting numeric tokens."""
    def __init__(self, raw):
        try:
            self.text = raw.decode('utf-8')
        except UnicodeError:
            raise Refusal('INVALID_JSON') from None
        self.pos = 0
        self.nodes = 0

    def whitespace(self):
        while self.pos < len(self.text) and self.text[self.pos] in ' \t\r\n':
            self.pos += 1

    def string(self):
        start = self.pos
        require(self.text[self.pos:self.pos + 1] == '"', 'INVALID_JSON')
        self.pos += 1
        while self.pos < len(self.text):
            char = self.text[self.pos]
            self.pos += 1
            if char == '\\':
                self.pos += 1
            elif char == '"':
                try:
                    return json.loads(self.text[start:self.pos])
                except (ValueError, UnicodeError):
                    break
        raise Refusal('INVALID_JSON')

    def value(self, depth=0):
        require(depth <= MAX_DEPTH, 'LIMIT_EXCEEDED')
        self.nodes += 1
        require(self.nodes <= MAX_NODES, 'LIMIT_EXCEEDED')
        self.whitespace()
        start = self.pos
        char = self.text[self.pos:self.pos + 1]
        if char in ('{', '['):
            kind = 'object' if char == '{' else 'array'
            close = '}' if char == '{' else ']'
            self.pos += 1
            self.whitespace()
            children, commas, keys = [], [], set()
            if self.text[self.pos:self.pos + 1] != close:
                while True:
                    key_start, key = self.pos, None
                    if kind == 'object':
                        key = self.string()
                        require(key not in keys, 'DUPLICATE_KEY')
                        keys.add(key)
                        self.whitespace()
                        require(self.text[self.pos:self.pos + 1] == ':', 'INVALID_JSON')
                        self.pos += 1
                    child = self.value(depth + 1)
                    children.append((key, key_start, child) if kind == 'object' else child)
                    self.whitespace()
                    if self.text[self.pos:self.pos + 1] != ',':
                        break
                    commas.append(self.pos)
                    self.pos += 1
                    self.whitespace()
            require(self.text[self.pos:self.pos + 1] == close, 'INVALID_JSON')
            self.pos += 1
            return Node(kind, start, self.pos, children, commas)
        if char == '"':
            self.string()
        elif char in ('t', 'f', 'n'):
            token = {'t': 'true', 'f': 'false', 'n': 'null'}[char]
            require(self.text.startswith(token, self.pos), 'INVALID_JSON')
            self.pos += len(token)
        else:
            match = NUMBER.match(self.text, self.pos)
            require(match is not None, 'INVALID_JSON')
            self.pos = match.end()
        return Node('scalar', start, self.pos)

    def parse(self):
        result = self.value()
        self.whitespace()
        require(self.pos == len(self.text), 'INVALID_JSON')
        return result


def target_paths(raw):
    # The same strict parser rejects duplicates/non-finite numbers and bounds
    # nesting before json.loads allocates the small target list.
    Parser(raw).parse()
    try:
        paths = json.loads(raw)
    except (ValueError, UnicodeError):
        raise Refusal('INVALID_TARGETS') from None
    require(isinstance(paths, list) and 0 < len(paths) <= MAX_TARGETS, 'INVALID_TARGETS')
    selected = set()
    for path in paths:
        require(isinstance(path, list) and 0 < len(path) <= MAX_DEPTH
                and all(isinstance(key, str) for key in path), 'INVALID_TARGETS')
        item = tuple(path)
        require(item not in selected, 'INVALID_TARGETS')
        selected.add(item)
    for path in selected:
        require(not any(path[:i] in selected for i in range(1, len(path))), 'OVERLAPPING_TARGETS')
    return selected


def shape(node, text, selected, path=()):
    if node.kind == 'scalar':
        return ('scalar', text[node.start:node.end])
    if node.kind == 'array':
        return ('array', tuple(shape(child, text, set()) for child in node.children))
    return ('object', tuple((key, shape(child, text, selected, path + (key,)))
                           for key, _, child in node.children if path + (key,) not in selected))


def candidate(raw, paths):
    parser = Parser(raw)
    root = parser.parse()
    removals = {}
    for path in paths:
        node = root
        for key in path[:-1]:
            require(node.kind == 'object', 'TARGET_NOT_OBJECT')
            matches = [child for name, _, child in node.children if name == key]
            require(bool(matches), 'TARGET_MISSING')
            node = matches[0]
        require(node.kind == 'object', 'TARGET_NOT_OBJECT')
        matches = [i for i, (key, _, _) in enumerate(node.children) if key == path[-1]]
        require(bool(matches), 'TARGET_MISSING')
        removals.setdefault(id(node), (node, set()))[1].add(matches[0])
    spans = []
    for node, indices in removals.values():
        ordered = sorted(indices)
        cursor = 0
        while cursor < len(ordered):
            first = last = ordered[cursor]
            cursor += 1
            while cursor < len(ordered) and ordered[cursor] == last + 1:
                last = ordered[cursor]
                cursor += 1
            if last < len(node.children) - 1:
                spans.append((node.children[first][1], node.commas[last] + 1))
            elif first:
                spans.append((node.commas[first - 1], node.children[last][2].end))
            else:
                spans.append((node.children[first][1], node.children[last][2].end))
    result, offset = [], 0
    for start, end in sorted(spans):
        require(start >= offset, 'VERIFY_FAILED')
        result.append(parser.text[offset:start])
        offset = end
    result.append(parser.text[offset:])
    output = ''.join(result).encode('utf-8')
    after = Parser(output)
    after_root = after.parse()
    require(shape(root, parser.text, paths) == shape(after_root, after.text, set()), 'VERIFY_FAILED')
    return output


def digest(raw):
    return hashlib.sha256(raw).hexdigest()


def write_candidate(path, data, source, verify):
    absolute(path)
    require(os.path.commonpath([path, source]) not in (path, source), 'OUTPUT_SOURCE_OVERLAP')
    parent, name = os.path.split(path)
    require(bool(name), 'INVALID_PATH')
    require(parent != os.path.dirname(source), 'OUTPUT_SOURCE_DIRECTORY')
    folder = directory(parent)
    created = known = None
    try:
        require(private(os.fstat(folder), 0o700), 'OUTPUT_NOT_PRIVATE')
        fd = os.open(name, os.O_WRONLY | os.O_CREAT | os.O_EXCL | os.O_NOFOLLOW, 0o600, dir_fd=folder)
        try:
            known = identity(os.fstat(fd))
            created = known[:2]
            os.fchmod(fd, 0o600)
            known = identity(os.fstat(fd))
            offset = 0
            while offset < len(data):
                count = os.write(fd, memoryview(data)[offset:])
                require(count > 0, 'IO_REFUSED')
                offset += count
                known = identity(os.fstat(fd))
            os.fsync(fd)
            written_id = known
            verify()
            observed, output_id = read_file(path, private_input=True)
            require(observed == data and output_id == written_id, 'OUTPUT_CHANGED')
            # Ensure the pathname still points to our exclusive output.
            require(identity(os.stat(name, dir_fd=folder, follow_symlinks=False))[:2] == created,
                    'OUTPUT_CHANGED')
        finally:
            os.close(fd)
    except BaseException:
        if created is not None:
            try:
                current = os.stat(name, dir_fd=folder, follow_symlinks=False)
                if identity(current)[:2] == created:
                    require(identity(current) == known, 'CLEANUP_INCOMPLETE')
                    os.unlink(name, dir_fd=folder)
            except FileNotFoundError:
                pass
            except OSError:
                raise Refusal('CLEANUP_INCOMPLETE') from None
        raise
    finally:
        os.close(folder)


class Arguments(argparse.ArgumentParser):
    def error(self, message):
        raise Refusal('INVALID_ARGUMENTS')


def main(argv=None):
    try:
        parser = Arguments(description=__doc__)
        parser.add_argument('--source', required=True)
        parser.add_argument('--targets', required=True)
        parser.add_argument('--expected-sha256')
        parser.add_argument('--candidate')
        args = parser.parse_args(argv)
        require(not args.candidate or args.expected_sha256 is not None, 'EXPECTED_HASH_REQUIRED')
        require(args.expected_sha256 is None or SHA.fullmatch(args.expected_sha256), 'INVALID_EXPECTED_HASH')
        source, source_id = read_file(args.source)
        target_data, target_id = read_file(args.targets, private_input=True)
        paths = target_paths(target_data)
        require(args.expected_sha256 is None or digest(source) == args.expected_sha256, 'SOURCE_HASH_MISMATCH')
        output = candidate(source, paths)

        def verify():
            fresh, fresh_id = read_file(args.source)
            targets, targets_id = read_file(args.targets, private_input=True)
            require(fresh == source and fresh_id == source_id
                    and targets == target_data and targets_id == target_id, 'SOURCE_CHANGED')
        verify()
        if args.candidate:
            write_candidate(args.candidate, output, args.source, verify)
        print(json.dumps({'status': 'CANDIDATE' if args.candidate else 'PREVIEW',
                          'targets': len(paths), 'source_sha256': digest(source),
                          'candidate_sha256': digest(output)}, sort_keys=True))
        return 0
    except Refusal as error:
        print(error.code, file=sys.stderr)
    except FileExistsError:
        print('OUTPUT_EXISTS', file=sys.stderr)
    except OSError:
        print('IO_REFUSED', file=sys.stderr)
    except (ValueError, TypeError, RecursionError, OverflowError):
        print('INVALID_INPUT', file=sys.stderr)
    except Exception:
        print('INTERNAL_ERROR', file=sys.stderr)
    return 2


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