mirror of
https://github.com/firestar5683/StarPilot.git
synced 2026-10-05 13:53:53 +08:00
Compatibility Shim
This commit is contained in:
@@ -0,0 +1,145 @@
|
||||
import fcntl
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
import stat
|
||||
import tempfile
|
||||
import uuid
|
||||
from contextlib import ExitStack
|
||||
|
||||
KEYS = ('CarParamsCache', 'CarParamsPersistent', 'CarParamsPrevRoute', 'CalibrationParams',
|
||||
'LiveParametersV2', 'LiveTorqueParameters', 'LiveDelay')
|
||||
PREFIX = b'SPCACHE'
|
||||
MAX_BYTES = 32 * 1024 * 1024
|
||||
|
||||
|
||||
def _sync(path):
|
||||
fd = os.open(path, os.O_RDONLY | os.O_DIRECTORY | os.O_NOFOLLOW)
|
||||
try:
|
||||
os.fsync(fd)
|
||||
finally:
|
||||
os.close(fd)
|
||||
|
||||
|
||||
def _read(path):
|
||||
fd = os.open(path, os.O_RDONLY | os.O_NOFOLLOW | os.O_NONBLOCK)
|
||||
with os.fdopen(fd, 'rb') as source:
|
||||
if not stat.S_ISREG(os.fstat(source.fileno()).st_mode):
|
||||
raise ValueError('Cache source must be a regular file')
|
||||
raw = source.read(MAX_BYTES + 1)
|
||||
if len(raw) > MAX_BYTES:
|
||||
raise ValueError('Cache exceeds recovery limit')
|
||||
return raw
|
||||
|
||||
|
||||
def _private(path):
|
||||
path.mkdir(mode=0o700, parents=True, exist_ok=True)
|
||||
info = path.lstat()
|
||||
if not stat.S_ISDIR(info.st_mode) or info.st_uid != os.getuid() or info.st_mode & 0o077:
|
||||
raise ValueError('Recovery storage must be private and owned')
|
||||
return path
|
||||
|
||||
|
||||
def _save(path, raw):
|
||||
fd, name = tempfile.mkstemp(prefix='.pending-', dir=path.parent)
|
||||
temporary = Path(name)
|
||||
try:
|
||||
with os.fdopen(fd, 'wb') as output:
|
||||
output.write(raw)
|
||||
output.flush()
|
||||
os.fsync(output.fileno())
|
||||
if _read(temporary) != raw:
|
||||
raise ValueError('Recovery temporary readback failed')
|
||||
try:
|
||||
os.link(temporary, path, follow_symlinks=False)
|
||||
except FileExistsError:
|
||||
info = path.lstat()
|
||||
if not stat.S_ISREG(info.st_mode) or info.st_uid != os.getuid() or info.st_mode & 0o077:
|
||||
raise ValueError('Unsafe recovery file') from None
|
||||
if _read(path) != raw:
|
||||
raise ValueError('Recovery archive readback failed')
|
||||
fd = os.open(path, os.O_RDONLY | os.O_NOFOLLOW)
|
||||
try:
|
||||
os.fsync(fd)
|
||||
finally:
|
||||
os.close(fd)
|
||||
finally:
|
||||
temporary.unlink(missing_ok=True)
|
||||
|
||||
|
||||
def _publish_handoff(namespace, target):
|
||||
name = '.starpilot-dom-handoff-' + hashlib.sha256(str(namespace).encode()).hexdigest() + '.json'
|
||||
marker = namespace.parent / name
|
||||
if marker.exists() or marker.is_symlink():
|
||||
info = marker.lstat()
|
||||
if not stat.S_ISREG(info.st_mode) or info.st_uid != os.getuid() or info.st_mode & 0o077:
|
||||
raise ValueError('Unsafe branch handoff marker')
|
||||
value = {'format': 'starpilot-dom-handoff', 'version': 1, 'namespace': str(namespace),
|
||||
'target': str(target), 'token': uuid.uuid4().hex}
|
||||
raw = json.dumps(value, sort_keys=True, separators=(',', ':')).encode()
|
||||
if len(raw) > 4096:
|
||||
raise ValueError('Branch handoff marker exceeds limit')
|
||||
fd, name = tempfile.mkstemp(prefix='.handoff-', dir=namespace.parent)
|
||||
temporary = Path(name)
|
||||
try:
|
||||
with os.fdopen(fd, 'wb') as output:
|
||||
output.write(raw)
|
||||
output.flush()
|
||||
os.fsync(output.fileno())
|
||||
if namespace.resolve(strict=True) != target:
|
||||
raise ValueError('Params namespace changed during handoff')
|
||||
os.replace(temporary, marker)
|
||||
_sync(namespace.parent)
|
||||
finally:
|
||||
temporary.unlink(missing_ok=True)
|
||||
|
||||
|
||||
def retire_foreign_caches(primary, secondary, archive):
|
||||
primary = Path(primary).absolute()
|
||||
namespaces = sorted({Path(path).absolute() for path in (primary, secondary)}, key=str)
|
||||
archive = Path(archive).absolute()
|
||||
with ExitStack() as locks:
|
||||
for root in sorted({path.parent for path in namespaces}, key=str):
|
||||
fd = os.open(root / '.lock', os.O_RDWR | os.O_CREAT | os.O_NOFOLLOW, 0o600)
|
||||
locks.callback(os.close, fd)
|
||||
if not stat.S_ISREG(os.fstat(fd).st_mode):
|
||||
raise ValueError('Invalid Params lock')
|
||||
fcntl.flock(fd, fcntl.LOCK_EX)
|
||||
targets = {path: path.resolve(strict=True) for path in namespaces}
|
||||
for path, target in targets.items():
|
||||
if target.parent != path.parent.resolve() or not target.is_dir():
|
||||
raise ValueError('Invalid Params namespace')
|
||||
if archive == target or target in archive.parents:
|
||||
raise ValueError('Recovery archive cannot be inside Params')
|
||||
originals = []
|
||||
for namespace in namespaces:
|
||||
for key in KEYS:
|
||||
try:
|
||||
raw = _read(targets[namespace] / key)
|
||||
except FileNotFoundError:
|
||||
continue
|
||||
if raw.startswith(PREFIX):
|
||||
originals.append((namespace, key, raw))
|
||||
if not originals:
|
||||
_publish_handoff(primary, targets[primary])
|
||||
return ()
|
||||
_private(archive)
|
||||
records = []
|
||||
for namespace, key, raw in originals:
|
||||
digest = hashlib.sha256(raw).hexdigest()
|
||||
_save(archive / digest, raw)
|
||||
records.append({'namespace': str(namespace), 'target': str(targets[namespace]), 'key': key,
|
||||
'sha256': digest, 'size': len(raw)})
|
||||
manifest = json.dumps(records, sort_keys=True, separators=(',', ':')).encode()
|
||||
_save(archive / (hashlib.sha256(manifest).hexdigest() + '.json'), manifest)
|
||||
_sync(archive)
|
||||
_sync(archive.parent)
|
||||
for namespace, key, raw in originals:
|
||||
if namespace.resolve(strict=True) != targets[namespace] or _read(targets[namespace] / key) != raw:
|
||||
raise ValueError('Cache source changed during recovery')
|
||||
for namespace, key, _raw in originals:
|
||||
(targets[namespace] / key).unlink()
|
||||
_sync(targets[namespace])
|
||||
_publish_handoff(primary, targets[primary])
|
||||
return tuple(records)
|
||||
@@ -0,0 +1,162 @@
|
||||
import pytest
|
||||
import re
|
||||
import importlib.util
|
||||
import hashlib
|
||||
import json
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
spec = importlib.util.spec_from_file_location('guard', Path(__file__).resolve().parents[1] / 'cache_compat.py')
|
||||
guard = importlib.util.module_from_spec(spec)
|
||||
spec.loader.exec_module(guard)
|
||||
|
||||
|
||||
class TestCacheCompatibility:
|
||||
@pytest.fixture(autouse=True)
|
||||
def setup(self, tmp_path):
|
||||
self.root = tmp_path
|
||||
self.primary = self.root / "params/d"
|
||||
self.secondary = self.root / 'params_cache/d'
|
||||
self.primary.parent.mkdir()
|
||||
self.secondary.parent.mkdir()
|
||||
target = self.primary.parent / '.tmp_dom'
|
||||
target.mkdir()
|
||||
self.primary.symlink_to(target)
|
||||
self.secondary.mkdir()
|
||||
self.archive = self.root / 'archive'
|
||||
|
||||
def run_guard(self):
|
||||
return guard.retire_foreign_caches(self.primary, self.secondary, self.archive)
|
||||
|
||||
def test_layers_preserve_valid_opposite_and_preferences(self):
|
||||
(self.primary / 'CarParamsPersistent').write_bytes(b'SPCACHE\x01foreign')
|
||||
(self.secondary / 'CarParamsPersistent').write_bytes(b'ordinary Dom raw')
|
||||
(self.primary / 'CalibrationParams').write_bytes(b'ordinary calibration')
|
||||
(self.secondary / 'LiveDelay').write_bytes(b'SPCACHE\tfuture')
|
||||
for root in (self.primary, self.secondary):
|
||||
(root / 'GithubSshKeys').write_bytes(b'auth')
|
||||
(root / 'IsMetric').write_bytes(b'1')
|
||||
records = self.run_guard()
|
||||
assert len(records) == 2
|
||||
assert not (self.primary / 'CarParamsPersistent').exists()
|
||||
assert not (self.secondary / 'LiveDelay').exists()
|
||||
assert (self.secondary / 'CarParamsPersistent').read_bytes() == b'ordinary Dom raw'
|
||||
assert (self.primary / 'CalibrationParams').read_bytes() == b'ordinary calibration'
|
||||
for record in records:
|
||||
raw = (self.archive / record['sha256']).read_bytes()
|
||||
assert raw.startswith(b'SPCACHE')
|
||||
assert (self.archive / record['sha256']).stat().st_mode & 511 == 384
|
||||
assert self.archive.stat().st_mode & 511 == 448
|
||||
assert self.run_guard() == ()
|
||||
for root in (self.primary, self.secondary):
|
||||
assert (root / 'GithubSshKeys').read_bytes() == b'auth'
|
||||
assert (root / 'IsMetric').read_bytes() == b'1'
|
||||
|
||||
def test_archive_failure_before_any_deletion_then_retry(self, monkeypatch):
|
||||
files = [self.primary / 'CarParamsPersistent', self.secondary / 'LiveDelay']
|
||||
for index, file in enumerate(files):
|
||||
file.write_bytes(b'SPCACHE' + bytes([index]))
|
||||
save = guard._save
|
||||
calls = []
|
||||
|
||||
def interrupted(path, raw):
|
||||
calls.append(path)
|
||||
if len(calls) == 2:
|
||||
raise OSError('disk full')
|
||||
save(path, raw)
|
||||
|
||||
with monkeypatch.context() as patches:
|
||||
patches.setattr(guard, '_save', interrupted)
|
||||
with pytest.raises(OSError):
|
||||
self.run_guard()
|
||||
assert all(file.exists() for file in files)
|
||||
assert list(self.primary.parent.glob('.starpilot-dom-handoff-*')) == []
|
||||
assert len(self.run_guard()) == 2
|
||||
|
||||
def test_interrupted_archive_fsync_retries(self, monkeypatch):
|
||||
file = self.primary / 'CarParamsPersistent'
|
||||
file.write_bytes(b'SPCACHEforeign')
|
||||
with monkeypatch.context() as patches:
|
||||
|
||||
def interrupted(*args, **kwargs):
|
||||
raise OSError('interrupted write')
|
||||
|
||||
patches.setattr(guard.os, 'fsync', interrupted)
|
||||
with pytest.raises(OSError):
|
||||
self.run_guard()
|
||||
assert file.exists()
|
||||
assert list(self.archive.iterdir()) == []
|
||||
assert len(self.run_guard()) == 1
|
||||
|
||||
def test_source_symlink_refused(self):
|
||||
outside = self.root / 'outside'
|
||||
outside.write_bytes(b'SPCACHEforeign')
|
||||
(self.primary / 'CarParamsPersistent').symlink_to(outside)
|
||||
with pytest.raises(OSError):
|
||||
self.run_guard()
|
||||
assert outside.read_bytes() == b'SPCACHEforeign'
|
||||
|
||||
def test_archive_symlink_refused(self):
|
||||
(self.primary / 'CarParamsPersistent').write_bytes(b'SPCACHEforeign')
|
||||
outside = self.root / 'outside'
|
||||
outside.mkdir()
|
||||
self.archive.symlink_to(outside)
|
||||
with pytest.raises(ValueError):
|
||||
self.run_guard()
|
||||
assert (self.primary / 'CarParamsPersistent').exists()
|
||||
|
||||
def test_nonforeign_malformed_bytes_are_preserved(self):
|
||||
(self.primary / 'CarParamsPersistent').write_bytes(b'bad nonword bytes')
|
||||
assert self.run_guard() == ()
|
||||
assert (self.primary / 'CarParamsPersistent').read_bytes() == b'bad nonword bytes'
|
||||
|
||||
def test_each_dom_start_marks_its_exact_namespace_even_without_foreign_caches(self):
|
||||
self.run_guard()
|
||||
name = '.starpilot-dom-handoff-' + hashlib.sha256(str(self.primary).encode()).hexdigest() + '.json'
|
||||
marker = self.primary.parent / name
|
||||
first = json.loads(marker.read_bytes())
|
||||
assert set(first) == {'format', 'version', 'namespace', 'target', 'token'}
|
||||
assert first['format'] == 'starpilot-dom-handoff'
|
||||
assert first['version'] == 1
|
||||
assert first['namespace'] == str(self.primary)
|
||||
assert first['target'] == str(self.primary.resolve())
|
||||
assert re.search('^[0-9a-f]{32}$', first['token'])
|
||||
assert marker.stat().st_mode & 511 == 384
|
||||
assert marker.stat().st_uid == os.getuid()
|
||||
self.run_guard()
|
||||
assert json.loads(marker.read_bytes())['token'] != first['token']
|
||||
|
||||
def test_handoff_failure_keeps_archives_and_can_retry_without_foreign_caches(self, monkeypatch):
|
||||
raw = b'SPCACHEforeign'
|
||||
(self.primary / 'CarParamsPersistent').write_bytes(raw)
|
||||
with monkeypatch.context() as patches:
|
||||
|
||||
def interrupted(*args, **kwargs):
|
||||
raise OSError('interrupted write')
|
||||
|
||||
patches.setattr(guard, '_publish_handoff', interrupted)
|
||||
with pytest.raises(OSError):
|
||||
self.run_guard()
|
||||
assert (self.archive / hashlib.sha256(raw).hexdigest()).read_bytes() == raw
|
||||
assert list(self.primary.parent.glob('.starpilot-dom-handoff-*')) == []
|
||||
assert self.run_guard() == ()
|
||||
assert len(list(self.primary.parent.glob('.starpilot-dom-handoff-*'))) == 1
|
||||
|
||||
def test_namespaces_have_independent_handoff_markers(self):
|
||||
self.run_guard()
|
||||
first = {path: path.read_bytes() for path in self.primary.parent.glob('.starpilot-dom-handoff-*')}
|
||||
named = self.primary.parent / 'test'
|
||||
named.mkdir()
|
||||
guard.retire_foreign_caches(named, self.secondary, self.archive)
|
||||
markers = list(self.primary.parent.glob('.starpilot-dom-handoff-*'))
|
||||
assert len(markers) == 2
|
||||
assert {path: path.read_bytes() for path in first} == first
|
||||
|
||||
def test_handoff_symlink_is_rejected(self):
|
||||
name = '.starpilot-dom-handoff-' + hashlib.sha256(str(self.primary).encode()).hexdigest() + '.json'
|
||||
outside = self.root / 'outside'
|
||||
outside.write_bytes(b'untouched')
|
||||
(self.primary.parent / name).symlink_to(outside)
|
||||
with pytest.raises(ValueError):
|
||||
self.run_guard()
|
||||
assert outside.read_bytes() == b'untouched'
|
||||
@@ -996,6 +996,10 @@ def manager_init() -> None:
|
||||
migrate_starpilot_param_renames(params, params_cache)
|
||||
last_timing = _log_boot_timing("manager_init", "param_renames", manager_init_start, last_timing)
|
||||
|
||||
from openpilot.starpilot.common.cache_compat import retire_foreign_caches
|
||||
retire_foreign_caches(params.get_param_path(), params_cache.get_param_path(),
|
||||
Path(cache_params_path).parent / 'dom-derived-cache-recovery')
|
||||
|
||||
params.clear_all(ParamKeyFlag.CLEAR_ON_MANAGER_START)
|
||||
params.clear_all(ParamKeyFlag.CLEAR_ON_ONROAD_TRANSITION)
|
||||
params.clear_all(ParamKeyFlag.CLEAR_ON_OFFROAD_TRANSITION)
|
||||
|
||||
Reference in New Issue
Block a user