Files
truf-server/app/worker_package.py
2026-09-30 20:30:56 +03:00

519 lines
22 KiB
Python

import hashlib
import json
import os
import platform as host_platform
import re
import stat
import sys
from lifecycle_authority import (
APPLICATION_IMPORT_SUFFIXES,
CODE_MANIFEST_SCHEMA,
GIT_MANIFEST_NAME,
REMOTE_WORKER_CODE_AUTHORITY_FILES,
TRUFFLEHOG_MANIFEST_NAME,
code_manifest_sha256,
verify_code_manifest,
)
from result_bundle import FORMAT_VERSION
from runtime_security import (
atomic_write_private_json,
canonical_path,
ensure_private_directory,
reject_reparse_components,
sha256_file,
)
WORKER_PACKAGE_SCHEMA = 3
PROTOCOL_VERSION = 2
PACKAGE_DETECTOR_POLICY = '@package/detector_policy'
MAX_WORKER_PACKAGE_MANIFEST_BYTES = 1024 * 1024
MAX_WORKER_PACKAGE_FILES = 512
MAX_WORKER_RUNTIME_TREE_FILES = 10000
MAX_WORKER_PACKAGE_CAPABILITIES = 16
_DIGEST = re.compile(r'^[a-f0-9]{64}$')
KNOWN_WORKER_PACKAGE_CAPABILITIES = frozenset({
('github', 'github', 'exact_git_v1'),
('gitlab', 'gitlab', 'exact_git_v1'),
('dockerhub', 'docker', 'docker_direct_v1'),
('huggingface', 'huggingface', 'huggingface_space_v1'),
})
DEFAULT_WORKER_PACKAGE_CAPABILITIES = (
('gitlab', 'gitlab', 'exact_git_v1'),
('dockerhub', 'docker', 'docker_direct_v1'),
('huggingface', 'huggingface', 'huggingface_space_v1'),
)
class WorkerPackageError(ValueError):
pass
def local_platform_tag():
machine = host_platform.machine().strip().lower().replace('amd64', 'x86_64')
system = 'windows' if sys.platform == 'win32' else 'linux' if sys.platform.startswith('linux') else ''
if not system or machine not in {'x86_64', 'aarch64', 'arm64'}:
raise WorkerPackageError('unsupported worker platform')
return f'{system}-{machine.replace("arm64", "aarch64")}'
def _relative_path(value, label):
value = str(value or '')
if (
not value or len(value) > 512 or '\\' in value or '\x00' in value
or value.startswith('/') or value.endswith('/')
):
raise WorkerPackageError(f'invalid worker package {label} path')
parts = value.split('/')
if any(part in ('', '.', '..') for part in parts):
raise WorkerPackageError(f'invalid worker package {label} path')
return '/'.join(parts)
def _entry(value, label):
if not isinstance(value, dict) or set(value) != {'path', 'sha256'}:
raise WorkerPackageError(f'invalid worker package {label} entry')
digest = str(value.get('sha256') or '')
if not _DIGEST.fullmatch(digest):
raise WorkerPackageError(f'invalid worker package {label} digest')
return {'path': _relative_path(value.get('path'), label), 'sha256': digest}
def _tree_entry(value, label):
if not isinstance(value, dict) or set(value) != {'path', 'sha256', 'file_count'}:
raise WorkerPackageError(f'invalid worker package {label} tree entry')
digest = str(value.get('sha256') or '')
try:
file_count = int(value.get('file_count'))
except (TypeError, ValueError, OverflowError):
raise WorkerPackageError(f'invalid worker package {label} tree file count') from None
if not _DIGEST.fullmatch(digest) or not 1 <= file_count <= MAX_WORKER_RUNTIME_TREE_FILES:
raise WorkerPackageError(f'invalid worker package {label} tree identity')
return {
'path': _relative_path(value.get('path'), label),
'sha256': digest,
'file_count': file_count,
}
def _validate_worker_application_files(names):
names = set(names)
required = set(REMOTE_WORKER_CODE_AUTHORITY_FILES)
if not required <= names:
raise WorkerPackageError('worker package file set is incomplete')
for name in names:
parts = name.lower().split('/')
if 'keycheckers' in parts or parts[-1] == 'keycheck_runner.py':
raise WorkerPackageError('worker package detailed keycheck code is forbidden')
application_code = {
name for name in names
if not name.startswith('dependencies/')
and name.lower().endswith(APPLICATION_IMPORT_SUFFIXES)
}
if application_code != required:
raise WorkerPackageError('worker package application code set is unsupported')
def normalize_worker_package_capabilities(value):
if (
not isinstance(value, list)
or not 1 <= len(value) <= MAX_WORKER_PACKAGE_CAPABILITIES
):
raise WorkerPackageError('worker package capabilities are invalid')
capabilities = []
seen = set()
for item in value:
if not isinstance(item, dict) or set(item) != {
'source', 'platform', 'planning_kind',
}:
raise WorkerPackageError('worker package capability shape is invalid')
if any(type(item[name]) is not str for name in item):
raise WorkerPackageError('worker package capability identity is invalid')
capability = (
item['source'], item['platform'], item['planning_kind'],
)
if capability not in KNOWN_WORKER_PACKAGE_CAPABILITIES:
raise WorkerPackageError('worker package capability is unsupported')
if capability in seen:
raise WorkerPackageError('worker package capability is duplicated')
seen.add(capability)
capabilities.append({
'source': capability[0],
'platform': capability[1],
'planning_kind': capability[2],
})
return sorted(
capabilities,
key=lambda item: (
item['source'], item['platform'], item['planning_kind'],
),
)
def normalize_worker_package_manifest(value):
if not isinstance(value, dict) or set(value) != {
'schema', 'protocol_version', 'bundle_format_version', 'platform_tag',
'capabilities', 'app_root', 'files', 'executables', 'assets',
'runtime_trees',
}:
raise WorkerPackageError('worker package manifest shape is invalid')
if type(value.get('schema')) is not int or value['schema'] != WORKER_PACKAGE_SCHEMA:
raise WorkerPackageError('worker package manifest schema is unsupported')
if (
type(value.get('protocol_version')) is not int
or value['protocol_version'] != PROTOCOL_VERSION
):
raise WorkerPackageError('worker package protocol version is unsupported')
if (
type(value.get('bundle_format_version')) is not int
or value['bundle_format_version'] != FORMAT_VERSION
):
raise WorkerPackageError('worker package bundle format is unsupported')
platform_tag = str(value.get('platform_tag') or '')
if not re.fullmatch(r'(windows|linux)-(x86_64|aarch64)', platform_tag):
raise WorkerPackageError('worker package platform is unsupported')
capabilities = normalize_worker_package_capabilities(value.get('capabilities'))
app_root = _relative_path(value.get('app_root'), 'application root')
files_value = value.get('files')
if not isinstance(files_value, dict) or not 1 <= len(files_value) <= MAX_WORKER_PACKAGE_FILES:
raise WorkerPackageError('worker package file set is invalid')
_validate_worker_application_files(files_value)
files = {}
for name in sorted(files_value):
normalized_name = _relative_path(name, 'file name')
if normalized_name != name:
raise WorkerPackageError('worker package file name is not canonical')
files[name] = _entry(files_value[name], f'file {name}')
executables_value = value.get('executables')
if not isinstance(executables_value, dict) or set(executables_value) != {
TRUFFLEHOG_MANIFEST_NAME, GIT_MANIFEST_NAME,
}:
raise WorkerPackageError('worker package executable set is incomplete')
executables = {
name: _entry(executables_value[name], f'executable {name}')
for name in (TRUFFLEHOG_MANIFEST_NAME, GIT_MANIFEST_NAME)
}
assets_value = value.get('assets')
if not isinstance(assets_value, dict) or set(assets_value) != {'detector_policy'}:
raise WorkerPackageError('worker package asset set is incomplete')
assets = {'detector_policy': _entry(assets_value['detector_policy'], 'detector policy')}
runtime_trees_value = value.get('runtime_trees')
required_trees = {'git', 'python'} if platform_tag.startswith('windows-') else {'git'}
if not isinstance(runtime_trees_value, dict) or set(runtime_trees_value) != required_trees:
raise WorkerPackageError('worker package runtime tree set is incomplete')
runtime_trees = {
name: _tree_entry(runtime_trees_value[name], f'{name} runtime')
for name in sorted(required_trees)
}
git_path = executables[GIT_MANIFEST_NAME]['path']
git_root = runtime_trees['git']['path'] + '/'
if not git_path.startswith(git_root):
raise WorkerPackageError('worker package Git executable escapes its runtime tree')
return {
'schema': WORKER_PACKAGE_SCHEMA,
'protocol_version': PROTOCOL_VERSION,
'bundle_format_version': FORMAT_VERSION,
'platform_tag': platform_tag,
'capabilities': capabilities,
'app_root': app_root,
'files': files,
'executables': executables,
'assets': assets,
'runtime_trees': runtime_trees,
}
def worker_package_manifest_sha256(manifest):
normalized = normalize_worker_package_manifest(manifest)
encoded = json.dumps(
normalized, ensure_ascii=True, sort_keys=True, separators=(',', ':'),
).encode('utf-8')
return hashlib.sha256(encoded).hexdigest()
def worker_package_build_compatibility(manifest):
normalized = normalize_worker_package_manifest(manifest)
return {
'protocol_version': normalized['protocol_version'],
'bundle_format_version': normalized['bundle_format_version'],
'platform_tag': normalized['platform_tag'],
'code_manifest_sha256': worker_package_manifest_sha256(normalized),
'detector_policy_sha256': normalized['assets']['detector_policy']['sha256'],
}
def load_worker_package_manifest_bytes(payload):
if type(payload) is not bytes:
raise WorkerPackageError('worker package manifest payload must be bytes')
if len(payload) > MAX_WORKER_PACKAGE_MANIFEST_BYTES:
raise WorkerPackageError('worker package manifest exceeds its byte bound')
try:
value = json.loads(payload.decode('utf-8', errors='strict'))
except (UnicodeDecodeError, json.JSONDecodeError) as exc:
raise WorkerPackageError('worker package manifest is invalid JSON') from exc
return normalize_worker_package_manifest(value)
def load_worker_package_manifest(path):
path = os.path.abspath(os.fspath(path))
reject_reparse_components(path)
flags = os.O_RDONLY
if hasattr(os, 'O_BINARY'):
flags |= os.O_BINARY
if hasattr(os, 'O_NOFOLLOW'):
flags |= os.O_NOFOLLOW
descriptor = os.open(path, flags)
with os.fdopen(descriptor, 'rb') as handle:
before = os.fstat(handle.fileno())
if not stat.S_ISREG(before.st_mode) or before.st_nlink != 1:
raise WorkerPackageError('worker package manifest is not a regular file')
payload = handle.read(MAX_WORKER_PACKAGE_MANIFEST_BYTES + 1)
after = os.fstat(handle.fileno())
reject_reparse_components(path)
current = os.stat(path, follow_symlinks=False)
identity = lambda value: (value.st_dev, value.st_ino)
if (
identity(before) != identity(after)
or identity(after) != identity(current)
or after.st_nlink != 1
or current.st_nlink != 1
or not stat.S_ISREG(current.st_mode)
or before.st_size != after.st_size
or after.st_size != current.st_size
or getattr(before, 'st_mtime_ns', None) != getattr(after, 'st_mtime_ns', None)
or getattr(after, 'st_mtime_ns', None) != getattr(current, 'st_mtime_ns', None)
or getattr(before, 'st_ctime_ns', None) != getattr(after, 'st_ctime_ns', None)
or (
os.name != 'nt'
and getattr(after, 'st_ctime_ns', None)
!= getattr(current, 'st_ctime_ns', None)
)
):
raise WorkerPackageError('worker package manifest changed while it was being read')
return load_worker_package_manifest_bytes(payload)
def _package_path(root, relative):
root = canonical_path(root)
candidate = canonical_path(os.path.join(root, *relative.split('/')))
try:
contained = os.path.commonpath((root, candidate)) == root and candidate != root
except ValueError:
contained = False
if not contained:
raise WorkerPackageError('worker package path escapes its root')
reject_reparse_components(candidate)
return candidate
def _manifest_entry_for_path(package_root, relative, label):
relative = _relative_path(relative, label)
path = _package_path(package_root, relative)
details = os.stat(path, follow_symlinks=False)
if not stat.S_ISREG(details.st_mode) or os.path.islink(path):
raise WorkerPackageError(f'worker package {label} is not a regular file')
return {'path': relative, 'sha256': sha256_file(path)}
def _application_package_files(package_root, app_root):
app_path = _package_path(package_root, app_root)
files = {}
def raise_walk_error(exc):
raise WorkerPackageError(f'unable to inspect worker package application root: {exc}') from exc
for current, directories, names in os.walk(
app_path, followlinks=False, onerror=raise_walk_error,
):
for name in directories:
path = os.path.join(current, name)
reject_reparse_components(path)
if os.path.islink(path) or name.lower() == '__pycache__':
raise WorkerPackageError('worker package application directory is unsupported')
for name in names:
path = os.path.join(current, name)
reject_reparse_components(path)
details = os.stat(path, follow_symlinks=False)
if not stat.S_ISREG(details.st_mode) or os.path.islink(path):
raise WorkerPackageError('worker package application file is not regular')
relative_name = os.path.relpath(path, app_path).replace(os.sep, '/')
relative_name = _relative_path(relative_name, 'file name')
files[relative_name] = _manifest_entry_for_path(
package_root, f'{app_root}/{relative_name}', f'file {relative_name}',
)
if len(files) > MAX_WORKER_PACKAGE_FILES:
raise WorkerPackageError('worker package file set exceeds its bound')
return files
def _runtime_tree_identity(package_root, relative, label):
relative = _relative_path(relative, label)
tree_root = _package_path(package_root, relative)
if not os.path.isdir(tree_root) or os.path.islink(tree_root):
raise WorkerPackageError(f'worker package {label} runtime tree is not a directory')
files = []
def raise_walk_error(exc):
raise WorkerPackageError(f'unable to inspect worker package {label} runtime tree: {exc}') from exc
for current, directories, names in os.walk(
tree_root, followlinks=False, onerror=raise_walk_error,
):
for name in directories:
path = os.path.join(current, name)
reject_reparse_components(path)
if os.path.islink(path):
raise WorkerPackageError(f'worker package {label} runtime tree contains a link')
for name in names:
path = os.path.join(current, name)
reject_reparse_components(path)
details = os.stat(path, follow_symlinks=False)
if not stat.S_ISREG(details.st_mode) or os.path.islink(path):
raise WorkerPackageError(f'worker package {label} runtime file is not regular')
name = _relative_path(
os.path.relpath(path, tree_root).replace(os.sep, '/'),
f'{label} runtime file',
)
files.append((name, sha256_file(path)))
if len(files) > MAX_WORKER_RUNTIME_TREE_FILES:
raise WorkerPackageError(f'worker package {label} runtime tree exceeds its bound')
if not files:
raise WorkerPackageError(f'worker package {label} runtime tree is empty')
digest = hashlib.sha256(b'truf-worker-runtime-tree-v1\0')
for name, file_digest in sorted(files):
digest.update(name.encode('utf-8', errors='strict'))
digest.update(b'\0')
digest.update(file_digest.encode('ascii'))
digest.update(b'\0')
return {'path': relative, 'sha256': digest.hexdigest(), 'file_count': len(files)}
def _runtime_tree_permissions_ready(package_root, entry):
tree_root = _package_path(package_root, entry['path'])
for current, directories, names in os.walk(tree_root, followlinks=False):
for path in [current, *(os.path.join(current, name) for name in directories + names)]:
details = os.stat(path, follow_symlinks=False)
if os.name == 'nt':
from runtime_security import private_directory_ready, private_file_ready
ready = private_directory_ready(path) if stat.S_ISDIR(details.st_mode) else private_file_ready(path)
if not ready:
return False
elif details.st_uid != 0 or details.st_mode & 0o022:
return False
return True
def _verify_runtime_tree(package_root, name, entry):
current = _runtime_tree_identity(package_root, entry['path'], name)
if current != entry:
raise WorkerPackageError(f'worker package {name} runtime tree drifted')
if not _runtime_tree_permissions_ready(package_root, entry):
raise WorkerPackageError(f'worker package {name} runtime tree is not trusted')
return _package_path(package_root, entry['path'])
def build_worker_package_manifest(
package_root, *, app_root='app', trufflehog_path, git_path,
detector_policy_path, git_root, capabilities, platform_tag=None,
python_root=None,
):
package_root = canonical_path(package_root)
reject_reparse_components(package_root)
if not os.path.isdir(package_root) or os.path.islink(package_root):
raise WorkerPackageError('worker package root is not a directory')
app_root = _relative_path(app_root, 'application root')
app_path = _package_path(package_root, app_root)
if not os.path.isdir(app_path) or os.path.islink(app_path):
raise WorkerPackageError('worker package application root is not a directory')
files = _application_package_files(package_root, app_root)
platform_tag = platform_tag or local_platform_tag()
runtime_trees = {'git': _runtime_tree_identity(package_root, git_root, 'git')}
if platform_tag.startswith('windows-'):
if not python_root:
raise WorkerPackageError('Windows worker package requires a Python runtime tree')
runtime_trees['python'] = _runtime_tree_identity(package_root, python_root, 'python')
elif python_root:
raise WorkerPackageError('Linux worker package must use its pinned image Python runtime')
manifest = {
'schema': WORKER_PACKAGE_SCHEMA,
'protocol_version': PROTOCOL_VERSION,
'bundle_format_version': FORMAT_VERSION,
'platform_tag': platform_tag,
'capabilities': list(capabilities),
'app_root': app_root,
'files': files,
'executables': {
TRUFFLEHOG_MANIFEST_NAME: _manifest_entry_for_path(
package_root, trufflehog_path, 'TruffleHog executable',
),
GIT_MANIFEST_NAME: _manifest_entry_for_path(
package_root, git_path, 'Git executable',
),
},
'assets': {
'detector_policy': _manifest_entry_for_path(
package_root, detector_policy_path, 'detector policy',
),
},
'runtime_trees': runtime_trees,
}
return normalize_worker_package_manifest(manifest)
def write_worker_package_manifest(path, manifest):
path = os.path.abspath(os.fspath(path))
ensure_private_directory(os.path.dirname(path), reject_reparse=True)
normalized = normalize_worker_package_manifest(manifest)
atomic_write_private_json(
path, normalized, max_bytes=MAX_WORKER_PACKAGE_MANIFEST_BYTES,
)
return path
def verify_worker_package(manifest_path):
manifest_path = os.path.abspath(os.fspath(manifest_path))
package_root = canonical_path(os.path.dirname(manifest_path))
manifest = load_worker_package_manifest(manifest_path)
if manifest['platform_tag'] != local_platform_tag():
raise WorkerPackageError('worker package does not match the local platform')
app_root = _package_path(package_root, manifest['app_root'])
files = {}
for name, entry in manifest['files'].items():
files[name] = {'path': _package_path(package_root, entry['path']), 'sha256': entry['sha256']}
executables = {
name: {'path': _package_path(package_root, entry['path']), 'sha256': entry['sha256']}
for name, entry in manifest['executables'].items()
}
policy = manifest['assets']['detector_policy']
policy_path = _package_path(package_root, policy['path'])
runtime_trees = {
name: _verify_runtime_tree(package_root, name, entry)
for name, entry in manifest['runtime_trees'].items()
}
code_manifest = {
'schema': CODE_MANIFEST_SCHEMA,
'root': app_root,
'files': files,
'executables': executables,
'assets': {policy_path: {'path': policy_path, 'sha256': policy['sha256']}},
}
verified = verify_code_manifest(
code_manifest,
require_private_acl=True,
required_names=REMOTE_WORKER_CODE_AUTHORITY_FILES,
external_names=(),
)
return {
'manifest': manifest,
'build_compatibility': worker_package_build_compatibility(manifest),
'code_manifest': verified,
'code_manifest_sha256': code_manifest_sha256(verified),
'trufflehog_path': executables[TRUFFLEHOG_MANIFEST_NAME]['path'],
'git_path': executables[GIT_MANIFEST_NAME]['path'],
'detector_policy_path': policy_path,
'runtime_trees': runtime_trees,
}