From 9263fd7c442edd20141150efbbe87acd385f08d2 Mon Sep 17 00:00:00 2001 From: firestar5683 <168790843+firestar5683@users.noreply.github.com> Date: Mon, 24 Aug 2026 21:29:41 -0500 Subject: [PATCH] Agnos 19.6.10 --- SConstruct | 45 +- launch_env.sh | 2 +- pyproject.toml | 5 + selfdrive/ui/SConscript | 30 +- selfdrive/ui/installer/installer.cc | 2 +- system/hardware/tici/agnos.json | 10 +- system/loggerd/tests/test_uploader.py | 6 +- system/loggerd/uploader.py | 3 + system/manager/manager.py | 5 - system/manager/process.py | 14 +- system/webrtc/device/video.py | 35 +- system/webrtc/tests/test_stream_session.py | 64 +- system/webrtc/tests/test_webrtcd.py | 79 +- system/webrtc/webrtcd.py | 399 +-- teleoprtc_repo/teleoprtc/builder.py | 29 +- teleoprtc_repo/teleoprtc/decoder.py | 50 + teleoprtc_repo/teleoprtc/info.py | 28 +- teleoprtc_repo/teleoprtc/stream.py | 455 ++-- teleoprtc_repo/teleoprtc/tracks.py | 64 +- tools/agnos/flash_desktop_system_to_comma.sh | 118 +- tools/agnos/patch_system_reset_image.py | 2245 +++++++---------- tools/agnos/test_patch_system_reset_image.py | 308 ++- tools/agnos/validate_agnos_runtime.sh | 90 + .../extract_sysroot_from_agnos.py | 1 + uv.lock | 39 + 25 files changed, 2199 insertions(+), 1927 deletions(-) create mode 100644 teleoprtc_repo/teleoprtc/decoder.py mode change 100644 => 100755 tools/agnos/patch_system_reset_image.py create mode 100755 tools/agnos/validate_agnos_runtime.sh diff --git a/SConstruct b/SConstruct index 7147ecdc1..3dd8c9337 100644 --- a/SConstruct +++ b/SConstruct @@ -1,4 +1,5 @@ import os +import importlib import shutil import subprocess import sys @@ -128,6 +129,34 @@ elif arch == "aarch64" and AGNOS: arch = "larch64" assert arch in ["larch64", "aarch64", "x86_64", "Darwin"] +# AGNOS 19.6 ships native dependencies as versioned Python packages. Link +# Cap'n Proto statically from that managed package so release binaries don't +# depend on the removed libcapnp-1.0.2.so system library. +try: + capnproto = importlib.import_module("capnproto") +except ModuleNotFoundError: + capnproto = None +try: + ffmpeg = importlib.import_module("ffmpeg") +except ModuleNotFoundError: + ffmpeg = None + +capnproto_include_dirs = [capnproto.INCLUDE_DIR] if capnproto is not None else [] +capnproto_lib_dirs = [capnproto.LIB_DIR] if capnproto is not None else [] +ffmpeg_include_dirs = [ffmpeg.INCLUDE_DIR] if ffmpeg is not None else [] +ffmpeg_lib_dirs = [ffmpeg.LIB_DIR] if ffmpeg is not None else [] + +# The managed native-dependency packages keep their tools inside the package +# instead of installing them into /usr/local/venv/bin. cereal invokes capnpc +# directly while SConscript files are evaluated, so make the packaged tools +# discoverable to both SCons actions and configure-time subprocesses. +dependency_bin_dirs = [ + package.BIN_DIR for package in (capnproto, ffmpeg) + if package is not None and os.path.isdir(package.BIN_DIR) +] +if dependency_bin_dirs: + os.environ["PATH"] = os.pathsep.join([*dependency_bin_dirs, os.environ["PATH"]]) + # Homebrew llvm can shadow Apple clang and break macOS SDK header resolution. # Use the system toolchain explicitly on macOS for reliable local builds. cc = '/usr/bin/clang' if arch == "Darwin" else 'clang' @@ -270,6 +299,8 @@ env = Environment( ] + cflags + ccflags, CPPPATH=cpppath + [ + capnproto_include_dirs, + ffmpeg_include_dirs, "#", "#third_party/acados/include", "#third_party/acados/include/blasfeo/include", @@ -288,11 +319,13 @@ env = Environment( RANLIB=ranlib, LINKFLAGS=ldflags, - RPATH=rpath, + RPATH=rpath + ffmpeg_lib_dirs, CFLAGS=["-std=gnu11"] + cflags, CXXFLAGS=["-std=c++1z"] + cxxflags, LIBPATH=libpath + [ + capnproto_lib_dirs, + ffmpeg_lib_dirs, "#msgq_repo", "#third_party", "#selfdrive/pandad", @@ -371,7 +404,15 @@ SConscript(['opendbc_repo/SConscript'], exports={'env': env_swaglog}) SConscript(['cereal/SConscript']) Import('socketmaster', 'msgq') -messaging = [socketmaster, msgq, 'capnp', 'kj',] +if capnproto is not None: + messaging = [ + socketmaster, + msgq, + File(os.path.join(capnproto.LIB_DIR, "libcapnp.a")), + File(os.path.join(capnproto.LIB_DIR, "libkj.a")), + ] +else: + messaging = [socketmaster, msgq, 'capnp', 'kj'] Export('messaging') diff --git a/launch_env.sh b/launch_env.sh index 5693bbf91..56b378f72 100755 --- a/launch_env.sh +++ b/launch_env.sh @@ -21,7 +21,7 @@ fi export QCOM_PRIORITY=12 if [ -z "$AGNOS_VERSION" ]; then - export AGNOS_VERSION="19.6.2" + export AGNOS_VERSION="19.6.10" fi if [ -z "$AGNOS_ACCEPTED_VERSIONS" ]; then diff --git a/pyproject.toml b/pyproject.toml index 8dce92875..ec6813d98 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -29,6 +29,11 @@ dependencies = [ "setuptools", "numpy >=2.0", + # AGNOS 19.6 native build dependencies + "comma-deps-capnproto; python_version >= '3.12'", + "comma-deps-ffmpeg; python_version >= '3.12'", + "libdatachannel-py>=2026.1.0.dev2; python_version >= '3.12'", + # body / webrtcd "aiohttp", "aiortc", diff --git a/selfdrive/ui/SConscript b/selfdrive/ui/SConscript index 6177692d9..0867df47c 100644 --- a/selfdrive/ui/SConscript +++ b/selfdrive/ui/SConscript @@ -1,15 +1,29 @@ +import importlib.util +import os +from pathlib import Path + Import('env', 'arch', 'common') if GetOption('extras') and arch != "Darwin": raylib_env = env.Clone() - raylib_env['LIBPATH'] += [f'#third_party/raylib/{arch}/'] raylib_env['LINKFLAGS'].append('-Wl,-strip-debug') - raylib_libs = common + ["raylib"] if arch == "larch64": - raylib_libs += ["GLESv2", "wayland-client", "wayland-egl", "EGL"] + tici_sysroot = os.environ.get("SP_TICI_SYSROOT", "").strip().rstrip("/") + if tici_sysroot: + raylib_dir = Path(tici_sysroot) / "usr/local/venv/lib/python3.12/site-packages/raylib/install" + else: + raylib_spec = importlib.util.find_spec("raylib") + if raylib_spec is None or raylib_spec.submodule_search_locations is None: + raise RuntimeError("The managed raylib package is required to build the comma installer") + raylib_dir = Path(raylib_spec.submodule_search_locations[0]) / "install" + + raylib_env['CPPPATH'] += [str(raylib_dir / "include")] + raylib_env['LIBPATH'] += [str(raylib_dir / "lib")] + raylib_libs = common + ["raylib_comma", "GLESv2", "EGL", "gbm", "drm"] else: - raylib_libs += ["GL"] + raylib_env['LIBPATH'] += [f'#third_party/raylib/{arch}/'] + raylib_libs = common + ["raylib", "GL"] release = "release3" installers = [ @@ -23,11 +37,15 @@ if GetOption('extras') and arch != "Darwin": "ld -r -b binary -o $TARGET $SOURCE") inter = raylib_env.Command("installer/inter_ttf.o", "installer/inter-ascii.ttf", "ld -r -b binary -o $TARGET $SOURCE") + inter_bold = raylib_env.Command("installer/inter_bold.o", "../assets/fonts/Inter-Bold.ttf", + "ld -r -b binary -o $TARGET $SOURCE") + inter_light = raylib_env.Command("installer/inter_light.o", "../assets/fonts/Inter-Light.ttf", + "ld -r -b binary -o $TARGET $SOURCE") for name, branch in installers: defines = {'BRANCH': f"'\"{branch}\"'"} if "internal" in name: defines['INTERNAL'] = "1" obj = raylib_env.Object(f"installer/installers/installer_{name}.o", ["installer/installer.cc"], CPPDEFINES=defines) - installer = raylib_env.Program(f"installer/installers/installer_{name}", [obj, cont, inter], LIBS=raylib_libs) - assert installer[0].get_size() < 1900*1e3, installer[0].get_size() + installer = raylib_env.Program(f"installer/installers/installer_{name}", [obj, cont, inter, inter_bold, inter_light], LIBS=raylib_libs) + assert installer[0].get_size() < 2500*1e3, installer[0].get_size() diff --git a/selfdrive/ui/installer/installer.cc b/selfdrive/ui/installer/installer.cc index 3e03ac0ff..b79e18329 100644 --- a/selfdrive/ui/installer/installer.cc +++ b/selfdrive/ui/installer/installer.cc @@ -6,7 +6,7 @@ #include "common/swaglog.h" #include "common/util.h" #include "system/hardware/hw.h" -#include "third_party/raylib/include/raylib.h" +#include "raylib.h" int freshClone(); int cachedFetch(const std::string &cache); diff --git a/system/hardware/tici/agnos.json b/system/hardware/tici/agnos.json index 3e2c65740..a93bddf7c 100644 --- a/system/hardware/tici/agnos.json +++ b/system/hardware/tici/agnos.json @@ -56,7 +56,7 @@ }, { "name": "boot", - "url": "https://www.dropbox.com/scl/fi/9l9io42qfx2shr9er5jqx/boot9.img.xz?rlkey=lbtxz862kxbvn3jn98ejitj8p&st=vr35tgz7&dl=1", + "url": "https://files.firestar.link/x/ugiq4cqx08q7/boot9.img.xz", "hash": "ab2eba0f96b2f48efa376330c3eb509158361adf3ad9c20f269ec92457aa841f", "hash_raw": "ab2eba0f96b2f48efa376330c3eb509158361adf3ad9c20f269ec92457aa841f", "size": 48343040, @@ -67,13 +67,13 @@ }, { "name": "system", - "url": "https://www.dropbox.com/scl/fi/pewhzpqzi3aewuiaffc6m/system10.img.xz?rlkey=olzrzulhs93zzghnjrskmdwxt&st=exnfk2oz&dl=1", - "hash": "ab395d4c963a908ab86709f1a6580a62dd24cb34ee71cf5fd5ec29d7d48d0e10", - "hash_raw": "ab395d4c963a908ab86709f1a6580a62dd24cb34ee71cf5fd5ec29d7d48d0e10", + "url": "https://files.firestar.link/x/07530yj1jd6a/system18.img.xz", + "hash": "01c84930849f9be2bdbad5e9a8dda3a6fd2be95e81d1b3556574b18346934f49", + "hash_raw": "01c84930849f9be2bdbad5e9a8dda3a6fd2be95e81d1b3556574b18346934f49", "size": 4718592000, "sparse": false, "full_check": false, "has_ab": true, - "ondevice_hash": "ab395d4c963a908ab86709f1a6580a62dd24cb34ee71cf5fd5ec29d7d48d0e10" + "ondevice_hash": "01c84930849f9be2bdbad5e9a8dda3a6fd2be95e81d1b3556574b18346934f49" } ] diff --git a/system/loggerd/tests/test_uploader.py b/system/loggerd/tests/test_uploader.py index 2e13d357a..aa3b13dc3 100644 --- a/system/loggerd/tests/test_uploader.py +++ b/system/loggerd/tests/test_uploader.py @@ -7,7 +7,7 @@ from pathlib import Path from openpilot.system.hardware.hw import Paths from openpilot.common.swaglog import cloudlog -from openpilot.system.loggerd.uploader import main, UPLOAD_ATTR_NAME, UPLOAD_ATTR_VALUE +from openpilot.system.loggerd.uploader import clear_locks, main, UPLOAD_ATTR_NAME, UPLOAD_ATTR_VALUE from openpilot.system.loggerd.xattr_cache import getxattr from openpilot.system.loggerd.tests.loggerd_tests_common import UploaderTestCase @@ -36,6 +36,10 @@ log_handler = FakeLogHandler() cloudlog.addHandler(log_handler) +def test_clear_locks_missing_root(tmp_path): + clear_locks(str(tmp_path / "missing")) + + class TestUploader(UploaderTestCase): def setup_method(self): super().setup_method() diff --git a/system/loggerd/uploader.py b/system/loggerd/uploader.py index 827dc8128..06ed867a9 100755 --- a/system/loggerd/uploader.py +++ b/system/loggerd/uploader.py @@ -63,6 +63,9 @@ def listdir_by_creation(d: str) -> list[str]: return [] def clear_locks(root: str) -> None: + if not os.path.isdir(root): + return + for logdir in os.listdir(root): path = os.path.join(root, logdir) try: diff --git a/system/manager/manager.py b/system/manager/manager.py index 431635fa1..034b2f4c4 100755 --- a/system/manager/manager.py +++ b/system/manager/manager.py @@ -1094,11 +1094,6 @@ def manager_init() -> None: device=HARDWARE.get_device_type()) last_timing = _log_boot_timing("manager_init", "logging_ready", manager_init_start, last_timing) - # preimport all processes - for p in managed_processes.values(): - p.prepare() - last_timing = _log_boot_timing("manager_init", "preimport_processes", manager_init_start, last_timing) - # StarPilot variables install_starpilot(build_metadata, params) last_timing = _log_boot_timing("manager_init", "install_starpilot", manager_init_start, last_timing) diff --git a/system/manager/process.py b/system/manager/process.py index 434c917b0..39e9d8291 100644 --- a/system/manager/process.py +++ b/system/manager/process.py @@ -634,19 +634,7 @@ class PythonProcess(ManagerProcess): self.launcher = launcher def prepare(self) -> None: - if self.enabled: - cloudlog.info(f"preimporting {self.module}") - start = time.monotonic() - try: - importlib.import_module(self.module) - finally: - line = f"SP_BOOT_TIMING preimport {self.name} module={self.module} +{time.monotonic() - start:.3f}s" - try: - with open(os.environ.get("SP_BOOT_TIMING_LOG", "/tmp/starpilot_boot_timing.log"), "a") as f: - f.write(line + "\n") - except OSError: - pass - cloudlog.warning(line) + pass def start(self) -> None: # In case we only tried a non blocking stop we need to stop it before restarting diff --git a/system/webrtc/device/video.py b/system/webrtc/device/video.py index fbe19c3dc..888f1739a 100644 --- a/system/webrtc/device/video.py +++ b/system/webrtc/device/video.py @@ -1,17 +1,19 @@ import asyncio +from dataclasses import dataclass import struct import time -import av from teleoprtc.tracks import TiciVideoStreamTrack -from aiortc.mediastreams import MediaStreamError from cereal import messaging -from openpilot.common.params import Params from openpilot.common.realtime import DT_MDL +from openpilot.common.params import Params +# v4l2 buffer flag marking an encoded keyframe (linux/videodev2.h) V4L2_BUF_FLAG_KEYFRAME = 0x8 + +# arbitrary 16-byte UUID identifying openpilot frame-timing SEI messages TIMING_SEI_UUID = bytes([ 0xa5, 0xe0, 0xc4, 0xa4, 0x5b, 0x6e, 0x4e, 0x1e, 0x9c, 0x7e, 0x12, 0x34, 0x56, 0x78, 0x9a, 0xbc, @@ -19,6 +21,15 @@ TIMING_SEI_UUID = bytes([ _SEI_PREFIX = b'\x00\x00\x00\x01\x06\x05\x30' + TIMING_SEI_UUID +@dataclass(frozen=True) +class EncodedVideoFrame: + data: bytes + pts: int + + def __bytes__(self) -> bytes: + return self.data + + class LiveStreamVideoStreamTrack(TiciVideoStreamTrack): camera_to_sock_mapping = { "driver": "livestreamDriverEncodeData", @@ -52,6 +63,9 @@ class LiveStreamVideoStreamTrack(TiciVideoStreamTrack): if not enabled: self._seen_keyframe = False + def request_keyframe(self) -> None: + self.params.put("LivestreamRequestKeyframe", True, block=False) + def _build_frame_data(self, msg) -> bytes: encode_data = getattr(msg, msg.which()) if not self.timing_sei_enabled: @@ -68,9 +82,7 @@ class LiveStreamVideoStreamTrack(TiciVideoStreamTrack): async def recv(self): while True: - if self.readyState != "live": - raise MediaStreamError - + # while video is disabled, pause here without returning if not self.video_enabled: await asyncio.sleep(0.005) continue @@ -79,18 +91,11 @@ class LiveStreamVideoStreamTrack(TiciVideoStreamTrack): if msg is not None: if not self._seen_keyframe and (getattr(msg, msg.which()).idx.flags & V4L2_BUF_FLAG_KEYFRAME): self._seen_keyframe = True - self.params.put("LivestreamRequestKeyframe", False) + self.params.put("LivestreamRequestKeyframe", False, block=False) break await asyncio.sleep(0.005) - packet = av.Packet(self._build_frame_data(msg)) - packet.time_base = self._time_base - self._pts = ((time.monotonic_ns() - self._t0_ns) * self._clock_rate) // 1_000_000_000 - packet.pts = self._pts self.log_debug("track sending frame %d", self._pts) - return packet - - def codec_preference(self) -> str | None: - return "H264" + return EncodedVideoFrame(self._build_frame_data(msg), self._pts) diff --git a/system/webrtc/tests/test_stream_session.py b/system/webrtc/tests/test_stream_session.py index f8316d203..7077e9877 100644 --- a/system/webrtc/tests/test_stream_session.py +++ b/system/webrtc/tests/test_stream_session.py @@ -1,20 +1,17 @@ import asyncio import json import time -# for aiortc and its dependencies -import warnings -warnings.filterwarnings("ignore", category=DeprecationWarning) -warnings.filterwarnings("ignore", category=RuntimeWarning) # TODO: remove this when google-crc32c publish a python3.12 wheel -from aiortc import RTCDataChannel -from aiortc.mediastreams import VIDEO_CLOCK_RATE, VIDEO_TIME_BASE import capnp -import pyaudio +import pytest + +pytest.importorskip("libdatachannel", reason="the upstream WebRTC backend requires Python 3.12") + from cereal import messaging, log +from teleoprtc.tracks import VIDEO_CLOCK_RATE from openpilot.system.webrtc.webrtcd import CerealOutgoingMessageProxy, CerealIncomingMessageProxy from openpilot.system.webrtc.device.video import LiveStreamVideoStreamTrack -from openpilot.system.webrtc.device.audio import AudioInputStreamTrack class TestStreamSession: @@ -33,40 +30,37 @@ class TestStreamSession: expected_dict = {"type": "customReservedRawData0", "logMonoTime": 123, "valid": True, "data": "test"} expected_json = json.dumps(expected_dict).encode() - channel = mocker.Mock(spec=RTCDataChannel) - mocked_submaster = messaging.SubMaster(["customReservedRawData0"]) - def mocked_update(t): - mocked_submaster.update_msgs(0, [test_msg]) + channel = mocker.Mock() + channel.is_open.return_value = True + proxy = CerealOutgoingMessageProxy(["customReservedRawData0"]) + + def mocked_update(_): + proxy.sm.update_msgs(0, [test_msg]) mocker.patch.object(messaging.SubMaster, "update", side_effect=mocked_update) - proxy = CerealOutgoingMessageProxy(["customReservedRawData0"]) - proxy.sm = mocked_submaster proxy.add_channel(channel) - proxy.update() channel.send.assert_called_once_with(expected_json) def test_incoming_proxy(self, mocker): tested_msgs = [ - {"type": "customReservedRawData0", "data": "test"}, # primitive - {"type": "can", "data": [{"address": 0, "dat": "", "src": 0}]}, # list - {"type": "testJoystick", "data": {"axes": [0, 0], "buttons": [False]}}, # dict + {"type": "customReservedRawData0", "data": "test"}, + {"type": "can", "data": [{"address": 0, "dat": "", "src": 0}]}, + {"type": "testJoystick", "data": {"axes": [0, 0], "buttons": [False]}}, ] mocked_pubmaster = mocker.MagicMock(spec=messaging.PubMaster) - proxy = CerealIncomingMessageProxy(mocked_pubmaster) for msg in tested_msgs: proxy.send(json.dumps(msg).encode()) mocked_pubmaster.send.assert_called_once() - mt, md = mocked_pubmaster.send.call_args.args - assert mt == msg["type"] - assert isinstance(md, capnp._DynamicStructBuilder) - assert hasattr(md, msg["type"]) - + msg_type, message = mocked_pubmaster.send.call_args.args + assert msg_type == msg["type"] + assert isinstance(message, capnp._DynamicStructBuilder) + assert hasattr(message, msg_type) mocked_pubmaster.reset_mock() def test_livestream_track(self, mocker): @@ -78,29 +72,11 @@ class TestStreamSession: track = LiveStreamVideoStreamTrack("driver") assert track.id.startswith("driver") - assert track.codec_preference() == "H264" for i in range(5): packet = self.loop.run_until_complete(track.recv()) - assert packet.time_base == VIDEO_TIME_BASE if i == 0: start_ns = time.monotonic_ns() start_pts = packet.pts - assert abs(i + packet.pts - (start_pts + (((time.monotonic_ns() - start_ns) * VIDEO_CLOCK_RATE) // 1_000_000_000))) < 450 #5ms - assert packet.size == 0 - - def test_input_audio_track(self, mocker): - packet_time, rate = 0.02, 16000 - sample_count = int(packet_time * rate) - mocked_stream = mocker.MagicMock(spec=pyaudio.Stream) - mocked_stream.read.return_value = b"\x00" * 2 * sample_count - - config = {"open.side_effect": lambda *args, **kwargs: mocked_stream} - mocker.patch("pyaudio.PyAudio", spec=True, **config) - track = AudioInputStreamTrack(audio_format=pyaudio.paInt16, packet_time=packet_time, rate=rate) - - for i in range(5): - frame = self.loop.run_until_complete(track.recv()) - assert frame.rate == rate - assert frame.samples == sample_count - assert frame.pts == i * sample_count + assert abs(i + packet.pts - (start_pts + (((time.monotonic_ns() - start_ns) * VIDEO_CLOCK_RATE) // 1_000_000_000))) < 450 + assert bytes(packet) == b"" diff --git a/system/webrtc/tests/test_webrtcd.py b/system/webrtc/tests/test_webrtcd.py index 23a7f6ddc..9fb6a42e5 100644 --- a/system/webrtc/tests/test_webrtcd.py +++ b/system/webrtc/tests/test_webrtcd.py @@ -1,65 +1,42 @@ -import pytest -import asyncio import json -# for aiortc and its dependencies -import warnings -warnings.filterwarnings("ignore", category=DeprecationWarning) -warnings.filterwarnings("ignore", category=RuntimeWarning) # TODO: remove this when google-crc32c publish a python3.12 wheel -from openpilot.system.webrtc.webrtcd import get_stream +import pytest -import aiortc -from teleoprtc import WebRTCOfferBuilder -from parameterized import parameterized_class +pytest.importorskip("libdatachannel", reason="the upstream WebRTC backend requires Python 3.12") + +from openpilot.system.webrtc.webrtcd import ServerState, handle_get_schema, handle_post_notify, on_shutdown -@parameterized_class(("in_services", "out_services"), [ - (["testJoystick"], ["carState"]), - ([], ["carState"]), - (["testJoystick"], []), - ([], []), -]) @pytest.mark.asyncio -class TestWebrtcdProc: - async def assertCompletesWithTimeout(self, awaitable, timeout=1): - try: - async with asyncio.timeout(timeout): - await awaitable - except TimeoutError: - pytest.fail("Timeout while waiting for awaitable to complete") +async def test_get_schema(): + status, body, content_type = await handle_get_schema(ServerState(), "carState") - async def test_webrtcd(self, mocker): - mock_request = mocker.MagicMock() - async def connect(offer): - body = {'sdp': offer.sdp, 'init_camera': offer.video[0], 'enabled': True, - 'bridge_services_in': self.in_services, 'bridge_services_out': self.out_services} - mock_request.json.side_effect = mocker.AsyncMock(return_value=body) - response = await get_stream(mock_request) - response_json = json.loads(response.text) - return aiortc.RTCSessionDescription(**response_json) + assert status == 200 + assert content_type.startswith("application/json") + assert "carState" in json.loads(body) - builder = WebRTCOfferBuilder(connect) - builder.offer_to_receive_video_stream("road") - builder.offer_to_receive_audio_stream() - if len(self.in_services) > 0 or len(self.out_services) > 0: - builder.add_messaging() - stream = builder.stream() +@pytest.mark.asyncio +async def test_get_schema_rejects_unknown_service(): + with pytest.raises(AssertionError, match="Invalid service name"): + await handle_get_schema(ServerState(), "notARealService") - await self.assertCompletesWithTimeout(stream.start()) - await self.assertCompletesWithTimeout(stream.wait_for_connection()) - assert stream.has_incoming_video_track("road") - assert stream.has_incoming_audio_track() - assert stream.has_messaging_channel() == (len(self.in_services) > 0 or len(self.out_services) > 0) +@pytest.mark.asyncio +async def test_notify_and_shutdown_active_stream(mocker): + state = ServerState() + session = mocker.MagicMock() + session.stop = mocker.AsyncMock() + state.streams["test"] = session - video_track, audio_track = stream.get_incoming_video_track("road"), stream.get_incoming_audio_track() - await self.assertCompletesWithTimeout(video_track.recv()) - await self.assertCompletesWithTimeout(audio_track.recv()) + status, body, content_type = await handle_post_notify(state, {"type": "ping"}) - await self.assertCompletesWithTimeout(stream.stop()) + assert (status, body) == (200, b"OK") + assert content_type.startswith("text/plain") + channel = session.stream.get_messaging_channel.return_value + channel.send.assert_called_once_with(json.dumps({"type": "ping"})) - # cleanup, very implementation specific, test may break if it changes - assert mock_request.app["streams"].__setitem__.called, "Implementation changed, please update this test" - _, session = mock_request.app["streams"].__setitem__.call_args.args - await self.assertCompletesWithTimeout(session.post_run_cleanup()) + await on_shutdown(state) + + session.stop.assert_awaited_once() + assert state.streams == {} diff --git a/system/webrtc/webrtcd.py b/system/webrtc/webrtcd.py index 837eda817..b3bacc81e 100644 --- a/system/webrtc/webrtcd.py +++ b/system/webrtc/webrtcd.py @@ -1,33 +1,29 @@ #!/usr/bin/env python3 from abc import abstractmethod +from collections.abc import Callable import os import socket import time +import capnp import argparse import asyncio import contextlib import json import uuid import logging -from typing import Any, TYPE_CHECKING - -# aiortc and its dependencies have lots of internal warnings :( -import warnings -warnings.filterwarnings("ignore", category=DeprecationWarning) -warnings.filterwarnings("ignore", category=RuntimeWarning) # TODO: remove this when google-crc32c publish a python3.12 wheel - -import capnp -from aiohttp import web -if TYPE_CHECKING: - from aiortc.rtcdatachannel import RTCDataChannel -import aioice.ice +import signal +import threading +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from urllib.parse import urlparse, parse_qs +from typing import Any from openpilot.system.webrtc.helpers import StreamRequestBody from openpilot.system.webrtc.schema import generate_field from openpilot.common.params import Params from cereal import messaging, log +SESSION_TIMEOUT_SECONDS = 300 # socket trick: route lookup for 8.8.8.8 (nothing is sent or actually connected to) # return the source interfaces IP which is the default interface of the device @@ -41,20 +37,8 @@ def _default_route_ip() -> str | None: finally: s.close() -# aioice patch: gather ICE candidates only on the default-route interface -_get_host_addresses = aioice.ice.get_host_addresses -def _primary_host_addresses(use_ipv4: bool, use_ipv6: bool) -> list[str]: - addresses = _get_host_addresses(use_ipv4, use_ipv6) - primary = _default_route_ip() - if primary not in addresses: - return addresses - return [primary, ] -aioice.ice.get_host_addresses = _primary_host_addresses - - class AsyncTaskRunner: def __init__(self): - self.is_running = False self.task = None self.logger = logging.getLogger("webrtcd") @@ -83,10 +67,10 @@ class CerealOutgoingMessageProxy(AsyncTaskRunner): super().__init__() self.services = list(services) self.sm = messaging.SubMaster(self.services) - self.channels: list[RTCDataChannel] = [] + self.channels = [] self._enabled = enabled - def add_channel(self, channel: 'RTCDataChannel'): + def add_channel(self, channel): self.channels.append(channel) def enable(self, enable: bool): @@ -115,20 +99,17 @@ class CerealOutgoingMessageProxy(AsyncTaskRunner): outgoing_msg = {"type": service, "logMonoTime": mono_time, "valid": valid, "data": msg_dict} encoded_msg = json.dumps(outgoing_msg).encode() for channel in self.channels: + if not channel.is_open(): + continue channel.send(encoded_msg) async def run(self): - from aiortc.exceptions import InvalidStateError - while True: if not self._enabled: await asyncio.sleep(0.01) continue try: self.update() - except InvalidStateError: - self.logger.warning("Cereal outgoing proxy invalid state (connection closed)") - break except Exception: self.logger.exception("Cereal outgoing proxy failure") await asyncio.sleep(0.01) @@ -169,17 +150,17 @@ class LivestreamBitrateController(AsyncTaskRunner): high_level = 0.1 # drop immediately med_level = 0.05 # drop after # of samples low_level = 0 # raise after # of samples - down_samples = 5 # 1s + down_samples = 5 param_name = "LivestreamEncoderBitrate" - def __init__(self, peer_connection: Any, params: Params, enabled: bool = True): + def __init__(self, get_stats: Callable[[], dict[str, Any]], params: Params, enabled: bool = True): super().__init__() - self.pc = peer_connection + self.get_stats = get_stats self.params = params self.level = 2 self._publish(self.bitrates[self.level]) - self.prev_lost, self.prev_sent = None, None + self.prev_stats: tuple[Any, ...] | None = None self.counter = 0 self.up_samples = 5 # 1s self._auto = True @@ -196,7 +177,7 @@ class LivestreamBitrateController(AsyncTaskRunner): if not self._auto: continue - loss_rate = await self._sample() + loss_rate = self._sample() if loss_rate is None: continue if loss_rate >= self.med_level and self.level > 0: @@ -213,22 +194,18 @@ class LivestreamBitrateController(AsyncTaskRunner): self.counter = 0 self._publish(self.bitrates[self.level]) - async def _sample(self) -> float | None: - report = await self.pc.getStats() - packets_lost = packets_sent = 0 - for s in report.values(): - if s.type == "remote-inbound-rtp": - packets_lost += s.packetsLost - elif s.type == "outbound-rtp": - packets_sent += s.packetsSent - - if self.prev_lost is None: - self.prev_lost, self.prev_sent = packets_lost, packets_sent + def _sample(self) -> float | None: + report = next(iter(self.get_stats().values()), None) + if report is None: return None - lost_delta = max(0, packets_lost - self.prev_lost) - sent_delta = max(0, packets_sent - self.prev_sent) - self.prev_lost, self.prev_sent = packets_lost, packets_sent - return lost_delta / sent_delta if sent_delta else 0.0 + + current = (report.ssrc, report.fraction_lost, report.packets_lost, report.highest_seq_no, report.jitter, report.lsr, report.dlsr) + if self.prev_stats == current: + return None + self.prev_stats = current + + loss_rate = report.fraction_lost / 256 + return loss_rate def _publish(self, bitrate: float): self.params.put(self.param_name, bitrate) @@ -244,48 +221,41 @@ class LivestreamBitrateController(AsyncTaskRunner): class StreamSession: shared_pub_master = DynamicPubMaster([]) - def __init__(self, body: StreamRequestBody, debug_mode: bool = False): - if debug_mode: - from aiortc.mediastreams import AudioStreamTrack, VideoStreamTrack - from aiortc.contrib.media import MediaBlackhole + def __init__(self, body: StreamRequestBody): from openpilot.system.webrtc.device.video import LiveStreamVideoStreamTrack - from openpilot.system.webrtc.device.audio import AudioInputStreamTrack, AudioOutputSpeaker from teleoprtc.builder import WebRTCAnswerBuilder - from teleoprtc.info import parse_info_from_offer self.identifier = str(uuid.uuid4()) self.params = Params() - builder = WebRTCAnswerBuilder(body.sdp) - config = parse_info_from_offer(body.sdp) + builder = WebRTCAnswerBuilder(body.sdp, bind_address=_default_route_ip()) self.enabled = body.enabled - self.video_track = LiveStreamVideoStreamTrack(body.init_camera, self.enabled) if not debug_mode else VideoStreamTrack() - builder.add_video_stream(body.init_camera, self.video_track) - if config.expected_audio_track: - builder.add_audio_stream(AudioInputStreamTrack() if not debug_mode else AudioStreamTrack()) - if config.incoming_audio_track: - self.audio_output_cls = AudioOutputSpeaker if not debug_mode else MediaBlackhole - builder.offer_to_receive_audio_stream() + self.video_tracks = [] + for camera in body.cameras: + track = LiveStreamVideoStreamTrack(camera, self.enabled) + self.video_tracks.append(track) + builder.add_video_stream(camera, track) self.stream = builder.stream() + self.is_body = "testJoystick" in body.bridge_services_in + self.incoming_bridge: CerealIncomingMessageProxy | None = None self.incoming_bridge_services = body.bridge_services_in self.outgoing_bridge: CerealOutgoingMessageProxy | None = None self.bitrate_controller: LivestreamBitrateController | None = None - self.audio_output: AudioOutputSpeaker | MediaBlackhole | None = None if len(body.bridge_services_in) > 0: self.incoming_bridge = CerealIncomingMessageProxy(self.shared_pub_master) if len(body.bridge_services_out) > 0: self.outgoing_bridge = CerealOutgoingMessageProxy(body.bridge_services_out, self.enabled) - self.bitrate_controller = LivestreamBitrateController(self.stream.peer_connection, self.params, self.enabled) + self.bitrate_controller = LivestreamBitrateController(self.stream.get_receiver_report_stats, self.params, self.enabled) self.run_task: asyncio.Task | None = None self._cleanup_lock = asyncio.Lock() self._cleanup_done = False self.logger = logging.getLogger("webrtcd") self.logger.info( - "New stream session (%s), init camera %s, video enabled %s, incoming services %s, outgoing services %s", - self.identifier, body.init_camera, body.enabled, body.bridge_services_in, body.bridge_services_out, + "New stream session (%s), video cameras %s, video enabled %s, incoming services %s, outgoing services %s", + self.identifier, [t.id for t in self.video_tracks], body.enabled, body.bridge_services_in, body.bridge_services_out, ) def start(self): @@ -310,16 +280,21 @@ class StreamSession: match msg_type: case "livestreamCameraSwitch": - self.video_track.switch_camera(payload["data"]["camera"]) + # only needed for 1 track stream + if len(self.video_tracks) == 1: + self.video_tracks[0].switch_camera(payload["data"]["camera"]) case "livestreamSettings": - self.bitrate_controller.set_quality(payload["data"]["quality"]) + if self.bitrate_controller is not None: + self.bitrate_controller.set_quality(payload["data"]["quality"]) case "livestreamVideoEnable": enabled = payload["data"]["enabled"] self.enabled = enabled - self.video_track.enable(enabled) + for track in self.video_tracks: + track.enable(enabled) if self.outgoing_bridge is not None: self.outgoing_bridge.enable(enabled) - self.bitrate_controller.enable(enabled) + if self.bitrate_controller is not None: + self.bitrate_controller.enable(enabled) if not enabled: self.params.put("LivestreamRequestKeyframe", True) case "clockSync": @@ -328,15 +303,29 @@ class StreamSession: }}) self.stream.get_messaging_channel().send(pong) case "enableTimingSei": - if hasattr(self.video_track, 'timing_sei_enabled'): - self.video_track.timing_sei_enabled = bool(payload["data"]["enabled"]) + for track in self.video_tracks: + track.timing_sei_enabled = bool(payload["data"]["enabled"]) case _: - if payload.get("type") not in self.incoming_bridge_services: + if msg_type not in self.incoming_bridge_services: return - self.incoming_bridge.send(message) + if self.incoming_bridge is not None: + self.incoming_bridge.send(message) except Exception: self.logger.exception("Cereal incoming proxy failure") + async def run_normal_session(self): + try: + await asyncio.wait_for(self.stream.wait_for_disconnection(), timeout=SESSION_TIMEOUT_SECONDS) + except TimeoutError: + self.logger.warning("Stream session (%s) timed out after %d s", self.identifier, SESSION_TIMEOUT_SECONDS) + try: + self.stream.get_messaging_channel().send(json.dumps({"type": "disconnect", "data": "Session timed out"})) + except Exception: + pass + + async def run_body_session(self): + await self.stream.wait_for_disconnection() + async def run(self): try: self.params.put("LivestreamRequestKeyframe", True) @@ -349,15 +338,14 @@ class StreamSession: channel = self.stream.get_messaging_channel() self.outgoing_bridge.add_channel(channel) self.outgoing_bridge.start() - if self.stream.has_incoming_audio_track(): - track = self.stream.get_incoming_audio_track(buffered=False) - self.audio_output = self.audio_output_cls() - self.audio_output.addTrack(track) - self.audio_output.start() - self.bitrate_controller.start() + if self.bitrate_controller is not None: + self.bitrate_controller.start() self.logger.info("Stream session (%s) connected", self.identifier) - await self.stream.wait_for_disconnection() + if self.is_body: + await self.run_body_session() + else: + await self.run_normal_session() self.logger.info("Stream session (%s) ended", self.identifier) except Exception: self.logger.exception("Stream session failure") @@ -370,39 +358,52 @@ class StreamSession: return self._cleanup_done = True self.params.put("LivestreamRequestKeyframe", False) - await self.bitrate_controller.stop() + if self.bitrate_controller is not None: + await self.bitrate_controller.stop() if self.outgoing_bridge is not None: await self.outgoing_bridge.stop() - if self.video_track is not None: - self.video_track.stop() - self.video_track = None - if self.audio_output is not None: - self.audio_output.stop() - self.audio_output = None + for track in self.video_tracks: + track.stop() + self.video_tracks.clear() await self.stream.stop() -def schedule_teardown(app): - # if nothing connects for 5 seconds, tear down livestreaming processes - h = app.get('teardown') - if h: - h.cancel() +class ServerState: + def __init__(self): + self.streams: dict[str, StreamSession] = {} + self.stream_lock = asyncio.Lock() + self.teardown: asyncio.TimerHandle | None = None + + +# if nothing connects for 5 seconds, tear down livestreaming processes +def schedule_teardown(state: ServerState): + if state.teardown is not None: + state.teardown.cancel() + def clear(): - if not app['streams']: - Params().put_bool("IsLiveStreaming", False) - app['teardown'] = asyncio.get_running_loop().call_later(5.0, clear) + if not state.streams: + Params().put_bool("IsLiveStreaming", False) + + state.teardown = asyncio.get_running_loop().call_later(5.0, clear) -async def get_stream(request: 'web.Request'): - stream_dict, debug_mode = request.app['streams'], request.app['debug'] - raw_body = await request.json() - body = StreamRequestBody(**raw_body) +def _json_response(obj: Any, status: int = 200) -> tuple[int, bytes, str]: + return (status, json.dumps(obj).encode(), "application/json; charset=utf-8") - async with request.app['stream_lock']: + +def _text_response(text: str, status: int = 200) -> tuple[int, bytes, str]: + return (status, text.encode(), "text/plain; charset=utf-8") + + +async def handle_get_stream(state: ServerState, raw_body: bytes) -> tuple[int, bytes, str]: + stream_dict = state.streams + body = StreamRequestBody(**json.loads(raw_body)) + + async with state.stream_lock: # don't remove existing connection on prewarm request enabled = any(s.run_task and not s.run_task.done() and s.enabled for s in stream_dict.values()) if enabled and not body.enabled: - return web.json_response({"error": "busy", "message": "someone else is connected."}) + return _json_response({"error": "busy", "message": "someone else is connected."}) for sid, s in list(stream_dict.items()): if s.run_task and not s.run_task.done(): @@ -414,10 +415,15 @@ async def get_stream(request: 'web.Request'): await s.stop() stream_dict.pop(sid, None) - session = StreamSession(body, debug_mode) + session = StreamSession(body) stream_dict[session.identifier] = session try: - answer = await session.get_answer() + answer = await asyncio.wait_for(session.get_answer(), timeout=30) + except TimeoutError: + await session.stop() + stream_dict.pop(session.identifier, None) + logging.getLogger("webrtcd").exception("Timed out creating stream answer") + raise except Exception: await session.stop() stream_dict.pop(session.identifier, None) @@ -427,94 +433,189 @@ async def get_stream(request: 'web.Request'): def remove_finished_session(_: asyncio.Task) -> None: stream_dict.pop(session.identifier, None) - schedule_teardown(request.app) + schedule_teardown(state) + session.run_task.add_done_callback(remove_finished_session) - return web.json_response({"sdp": answer.sdp, "type": answer.type}) + return _json_response({"sdp": answer.sdp, "type": answer.type}) -async def get_schema(request: 'web.Request'): - services = request.query.get("services", "").split(",") +async def handle_get_schema(state: ServerState, services_param: str) -> tuple[int, bytes, str]: + services = services_param.split(",") services = [s for s in services if s] assert all(s in log.Event.schema.fields and not s.endswith("DEPRECATED") for s in services), "Invalid service name" schema_dict = {s: generate_field(log.Event.schema.fields[s]) for s in services} - return web.json_response(schema_dict) + return _json_response(schema_dict) -async def post_notify(request: 'web.Request'): - try: - payload = await request.json() - except Exception as e: - raise web.HTTPBadRequest(text="Invalid JSON") from e - - for session in list(request.app.get('streams', {}).values()): +async def handle_post_notify(state: ServerState, payload: Any) -> tuple[int, bytes, str]: + for session in list(state.streams.values()): try: ch = session.stream.get_messaging_channel() ch.send(json.dumps(payload)) except Exception: continue - return web.Response(status=200, text="OK") + return _text_response("OK") -async def on_shutdown(app: 'web.Application'): - for session in list(app['streams'].values()): +async def on_shutdown(state: ServerState): + for session in list(state.streams.values()): try: ch = session.stream.get_messaging_channel() ch.send(json.dumps({"type": "disconnect", "data": "device streaming has been stopped."})) except Exception: pass await session.stop() - del app['streams'] + state.streams.clear() -@web.middleware -async def error_middleware(request: 'web.Request', handler): - try: - return await handler(request) - except Exception as e: - logging.getLogger("webrtcd").exception("Unhandled error handling %s", request.path) - return web.json_response({"error": "exception", "message": f"{type(e).__name__}: {e}"}, status=500) +class WebrtcdHandler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + # path -> allowed methods (aiohttp registered POST /stream, POST /notify, GET /schema + its auto HEAD) + _routes = { + "/schema": ("GET", "HEAD"), + "/stream": ("POST",), + "/notify": ("POST",), + } + + def _send(self, status: int, body: bytes, content_type: str) -> None: + self.send_response(status) + self.send_header("Content-Type", content_type) + self.send_header("Content-Length", str(len(body))) + self.end_headers() + if self.command != "HEAD": + self.wfile.write(body) + + def _read_body(self) -> bytes: + length = int(self.headers.get("Content-Length", 0)) + return self.rfile.read(length) if length else b"" + + def _run(self, coro) -> tuple[int, bytes, str]: + return asyncio.run_coroutine_threadsafe(coro, self.server.loop).result() + + def _dispatch_request(self) -> None: + parsed = urlparse(self.path) + allowed = self._routes.get(parsed.path) + + try: + if allowed is None: + result = _json_response({"error": "not found"}, status=404) + elif self.command not in allowed: + result = _json_response({"error": "method not allowed"}, status=405) + elif parsed.path == "/schema": + services = parse_qs(parsed.query).get("services", [""])[0] + result = self._run(handle_get_schema(self.server.state, services)) + elif parsed.path == "/stream": + result = self._run(handle_get_stream(self.server.state, self._read_body())) + else: # /notify + try: + payload = json.loads(self._read_body()) + except Exception: + result = _json_response({"error": "bad request"}, status=400) + else: + result = self._run(handle_post_notify(self.server.state, payload)) + except Exception as e: + logging.getLogger("webrtcd").exception("Unhandled error handling %s", self.path) + result = _json_response({"error": "exception", "message": f"{type(e).__name__}: {e}"}, status=500) + + self._send(*result) + + def do_GET(self) -> None: + self._dispatch_request() + + def do_HEAD(self) -> None: + self._dispatch_request() + + def do_POST(self) -> None: + self._dispatch_request() + + def do_PUT(self) -> None: + self._dispatch_request() + + def do_DELETE(self) -> None: + self._dispatch_request() + + def do_PATCH(self) -> None: + self._dispatch_request() + + def do_OPTIONS(self) -> None: + self._dispatch_request() + + def log_message(self, format: str, *args: object) -> None: # noqa: A002 # stdlib override + # silence default access logging; errors are logged explicitly in _dispatch_request + pass -def prewarm_stream_session_imports(debug_mode: bool = False) -> None: - if debug_mode: - from aiortc.mediastreams import VideoStreamTrack - assert VideoStreamTrack +class WebrtcdHTTPServer(ThreadingHTTPServer): + daemon_threads = True + allow_reuse_address = True + state: ServerState + loop: asyncio.AbstractEventLoop + + +async def _shutdown(server: WebrtcdHTTPServer, state: ServerState, loop: asyncio.AbstractEventLoop) -> None: + # stop accepting new HTTP connections (blocks until serve_forever returns, so + # run it off the loop) then tear down active stream sessions. + await loop.run_in_executor(None, server.shutdown) + await on_shutdown(state) + loop.stop() + + +def prewarm_stream_session_imports() -> None: from openpilot.system.webrtc.device.video import LiveStreamVideoStreamTrack from teleoprtc.builder import WebRTCAnswerBuilder assert LiveStreamVideoStreamTrack assert WebRTCAnswerBuilder -def webrtcd_thread(host: str, port: int, debug: bool): - logging.basicConfig(level=logging.CRITICAL, handlers=[logging.StreamHandler()]) +def webrtcd_thread(host: str, port: int): + logging.basicConfig(level=logging.INFO, handlers=[logging.StreamHandler()]) prewarm_start = time.monotonic() - prewarm_stream_session_imports(debug) + prewarm_stream_session_imports() prewarm_end = time.monotonic() logging.getLogger("webrtcd").info(f"webrtc prewarm finished in {(prewarm_end - prewarm_start) * 1000} ms") - app = web.Application(middlewares=[error_middleware]) + loop = asyncio.new_event_loop() + asyncio.set_event_loop(loop) + state = ServerState() - app['streams'] = dict() - app['stream_lock'] = asyncio.Lock() - app['debug'] = debug - app.on_shutdown.append(on_shutdown) - app.router.add_post("/stream", get_stream) - app.router.add_post("/notify", post_notify) - app.router.add_get("/schema", get_schema) + server = WebrtcdHTTPServer((host, port), WebrtcdHandler) + server.state = state + server.loop = loop - web.run_app(app, host=host, port=port) + # serve HTTP on a daemon thread so the asyncio loop can own the main thread + http_thread = threading.Thread(target=server.serve_forever, name="webrtcd-http", daemon=True) + http_thread.start() + + shutting_down = False + shutdown_task = None + + def request_shutdown() -> None: + nonlocal shutting_down, shutdown_task + if shutting_down: + return + shutting_down = True + shutdown_task = loop.create_task(_shutdown(server, state, loop)) + + for sig in (signal.SIGINT, signal.SIGTERM): + loop.add_signal_handler(sig, request_shutdown) + + try: + loop.run_forever() + finally: + server.server_close() + loop.close() def main(): parser = argparse.ArgumentParser(description="WebRTC daemon") parser.add_argument("--host", type=str, default="0.0.0.0", help="Host to listen on") parser.add_argument("--port", type=int, default=5001, help="Port to listen on") - parser.add_argument("--debug", action="store_true", help="Enable debug mode") args = parser.parse_args() - webrtcd_thread(args.host, args.port, args.debug) + webrtcd_thread(args.host, args.port) if __name__=="__main__": diff --git a/teleoprtc_repo/teleoprtc/builder.py b/teleoprtc_repo/teleoprtc/builder.py index cb9f5d057..cc6565a47 100644 --- a/teleoprtc_repo/teleoprtc/builder.py +++ b/teleoprtc_repo/teleoprtc/builder.py @@ -1,9 +1,7 @@ import abc -from typing import Dict, List +from typing import Dict, List, Optional -import aiortc - -from teleoprtc.stream import WebRTCBaseStream, WebRTCOfferStream, WebRTCAnswerStream, ConnectionProvider +from teleoprtc.stream import RTCSessionDescription, WebRTCBaseStream, WebRTCOfferStream, WebRTCAnswerStream, ConnectionProvider from teleoprtc.tracks import TiciVideoStreamTrack, TiciTrackWrapper @@ -14,11 +12,12 @@ class WebRTCStreamBuilder(abc.ABC): class WebRTCOfferBuilder(WebRTCStreamBuilder): - def __init__(self, connection_provider: ConnectionProvider): + def __init__(self, connection_provider: ConnectionProvider, bind_address: Optional[str] = None): self.connection_provider = connection_provider + self.bind_address = bind_address self.requested_camera_types: List[str] = [] self.requested_audio = False - self.audio_tracks: List[aiortc.MediaStreamTrack] = [] + self.audio_tracks: List[object] = [] self.messaging_enabled = False def offer_to_receive_video_stream(self, camera_type: str): @@ -28,7 +27,7 @@ class WebRTCOfferBuilder(WebRTCStreamBuilder): def offer_to_receive_audio_stream(self): self.requested_audio = True - def add_audio_stream(self, track: aiortc.MediaStreamTrack): + def add_audio_stream(self, track: object): assert len(self.audio_tracks) == 0 self.audio_tracks = [track] @@ -43,32 +42,34 @@ class WebRTCOfferBuilder(WebRTCStreamBuilder): video_producer_tracks=[], audio_producer_tracks=self.audio_tracks, should_add_data_channel=self.messaging_enabled, + bind_address=self.bind_address, ) class WebRTCAnswerBuilder(WebRTCStreamBuilder): - def __init__(self, offer_sdp: str): + def __init__(self, offer_sdp: str, bind_address: Optional[str] = None): self.offer_sdp = offer_sdp - self.video_tracks: Dict[str, aiortc.MediaStreamTrack] = dict() + self.bind_address = bind_address + self.video_tracks: Dict[str, TiciVideoStreamTrack] = {} self.requested_audio = False - self.audio_tracks: List[aiortc.MediaStreamTrack] = [] + self.audio_tracks: List[object] = [] def offer_to_receive_audio_stream(self): self.requested_audio = True - def add_video_stream(self, camera_type: str, track: aiortc.MediaStreamTrack): + def add_video_stream(self, camera_type: str, track: object): assert camera_type not in self.video_tracks assert camera_type in ["driver", "wideRoad", "road"] if not isinstance(track, TiciVideoStreamTrack): track = TiciTrackWrapper(camera_type, track) self.video_tracks[camera_type] = track - def add_audio_stream(self, track: aiortc.MediaStreamTrack): + def add_audio_stream(self, track: object): assert len(self.audio_tracks) == 0 self.audio_tracks = [track] def stream(self) -> WebRTCBaseStream: - description = aiortc.RTCSessionDescription(sdp=self.offer_sdp, type="offer") + description = RTCSessionDescription(sdp=self.offer_sdp, type="offer") return WebRTCAnswerStream( description, consumed_camera_types=[], @@ -76,5 +77,5 @@ class WebRTCAnswerBuilder(WebRTCStreamBuilder): video_producer_tracks=list(self.video_tracks.values()), audio_producer_tracks=self.audio_tracks, should_add_data_channel=False, + bind_address=self.bind_address, ) - diff --git a/teleoprtc_repo/teleoprtc/decoder.py b/teleoprtc_repo/teleoprtc/decoder.py new file mode 100644 index 000000000..2a2cb96e0 --- /dev/null +++ b/teleoprtc_repo/teleoprtc/decoder.py @@ -0,0 +1,50 @@ +import dataclasses +import struct +from typing import List + + +@dataclasses.dataclass(frozen=True) +class RtcpReceiverReport: + ssrc: int + fraction_lost: int + packets_lost: int + highest_seq_no: int + jitter: int + lsr: int + dlsr: int + + +def _decode_receiver_reports(message: bytes) -> List[RtcpReceiverReport]: + reports: List[RtcpReceiverReport] = [] + offset = 0 + + while offset + 4 <= len(message): + flags, packet_type, length_words = struct.unpack_from("!BBH", message, offset) + packet_end = offset + (length_words + 1) * 4 + if flags >> 6 != 2 or packet_end > len(message): + break + + report_count = flags & 0x1F + if packet_type == 200: # Sender Report + report_offset = offset + 28 + elif packet_type == 201: # Receiver Report + report_offset = offset + 8 + else: + offset = packet_end + continue + + if report_offset + report_count * 24 > packet_end: + break + + for i in range(report_count): + block_offset = report_offset + i * 24 + ssrc, loss, highest_seq_no, jitter, lsr, dlsr = struct.unpack_from("!IIIIII", message, block_offset) + fraction_lost = loss >> 24 + packets_lost = loss & 0xFFFFFF + if packets_lost & 0x800000: + packets_lost -= 1 << 24 + reports.append(RtcpReceiverReport(ssrc, fraction_lost, packets_lost, highest_seq_no, jitter, lsr, dlsr)) + + offset = packet_end + + return reports diff --git a/teleoprtc_repo/teleoprtc/info.py b/teleoprtc_repo/teleoprtc/info.py index 537b71242..191d177cd 100644 --- a/teleoprtc_repo/teleoprtc/info.py +++ b/teleoprtc_repo/teleoprtc/info.py @@ -1,6 +1,6 @@ import dataclasses -import aiortc +from libdatachannel import Description @dataclasses.dataclass @@ -15,13 +15,23 @@ def parse_info_from_offer(sdp: str) -> StreamingMediaInfo: """ helper function to parse info about outgoing and incoming streams from an offer sdp """ - desc = aiortc.sdp.SessionDescription.parse(sdp) - audio_tracks = [m for m in desc.media if m.kind == "audio"] - video_tracks = [m for m in desc.media if m.kind == "video" and m.direction in ["recvonly", "sendrecv"]] - application_tracks = [m for m in desc.media if m.kind == "application"] - has_incoming_audio_track = next((t for t in audio_tracks if t.direction in ["sendonly", "sendrecv"]), None) is not None - has_incoming_datachannel = len(application_tracks) > 0 - expects_outgoing_audio_track = next((t for t in audio_tracks if t.direction in ["recvonly", "sendrecv"]), None) is not None + desc = Description(sdp, Description.Type.Offer) + n_video = 0 + expected_audio_track = False + incoming_audio_track = False + incoming_datachannel = desc.has_application() - return StreamingMediaInfo(len(video_tracks), expects_outgoing_audio_track, has_incoming_audio_track, has_incoming_datachannel) + for i in range(desc.media_count()): + media = desc.media(i) + if media is None: + continue + direction = media.direction() + if media.type() == "video" and direction in (Description.Direction.RecvOnly, Description.Direction.SendRecv): + n_video += 1 + elif media.type() == "audio": + if direction in (Description.Direction.RecvOnly, Description.Direction.SendRecv): + expected_audio_track = True + if direction in (Description.Direction.SendOnly, Description.Direction.SendRecv): + incoming_audio_track = True + return StreamingMediaInfo(n_video, expected_audio_track, incoming_audio_track, incoming_datachannel) diff --git a/teleoprtc_repo/teleoprtc/stream.py b/teleoprtc_repo/teleoprtc/stream.py index 208bd6965..cc4f0092d 100644 --- a/teleoprtc_repo/teleoprtc/stream.py +++ b/teleoprtc_repo/teleoprtc/stream.py @@ -1,13 +1,29 @@ import abc import asyncio +import contextlib import dataclasses import logging -from typing import Any, Awaitable, Callable, Dict, List, Optional +import random +from typing import Any, Awaitable, Callable, Dict, List, Optional, Tuple, Union -import aiortc -from aiortc.contrib.media import MediaRelay +from libdatachannel import ( + Configuration, + DataChannel, + Description, + FrameInfo, + H264RtpPacketizer, + IceServer, + NalUnit, + PeerConnection, + PliHandler, + RtcpNackResponder, + RtcpSrReporter, + RtpPacketizationConfig, + Track, +) -from teleoprtc.tracks import parse_video_track_id +from teleoprtc.decoder import RtcpReceiverReport, _decode_receiver_reports +from teleoprtc.tracks import TiciVideoStreamTrack, parse_video_track_id @dataclasses.dataclass @@ -16,138 +32,238 @@ class StreamingOffer: video: List[str] -ConnectionProvider = Callable[[StreamingOffer], Awaitable[aiortc.RTCSessionDescription]] -MessageHandler = Callable[[bytes], Awaitable[None]] +@dataclasses.dataclass +class RTCSessionDescription: + sdp: str + type: str + + +ConnectionProvider = Callable[[StreamingOffer], Awaitable[RTCSessionDescription]] +MessageHandler = Callable[[Union[bytes, str]], None] class WebRTCBaseStream(abc.ABC): + # destorying wrapper on close can cause deadlock + # TODO: upstream a fix to this + _retained_messaging_channels: List[DataChannel] = [] + _retain_messaging_channel_on_close = False + def __init__(self, consumed_camera_types: List[str], consume_audio: bool, - video_producer_tracks: List[aiortc.MediaStreamTrack], - audio_producer_tracks: List[aiortc.MediaStreamTrack], - should_add_data_channel: bool): - self.peer_connection = aiortc.RTCPeerConnection() - self.media_relay = MediaRelay() + video_producer_tracks: List[TiciVideoStreamTrack], + audio_producer_tracks: List[Any], + should_add_data_channel: bool, + bind_address: Optional[str] = None): + config = Configuration() + config.force_media_transport = True + config.disable_auto_negotiation = True + config.ice_servers = [IceServer("stun:stun.l.google.com:19302")] + if bind_address is not None: + config.bind_address = bind_address + + self.peer_connection = PeerConnection(config) self.expected_incoming_camera_types = consumed_camera_types self.expected_incoming_audio = consume_audio self.expected_number_of_incoming_media: Optional[int] = None - self.incoming_camera_tracks: Dict[str, aiortc.MediaStreamTrack] = dict() - self.incoming_audio_tracks: List[aiortc.MediaStreamTrack] = [] - self.outgoing_video_tracks: List[aiortc.MediaStreamTrack] = video_producer_tracks - self.outgoing_audio_tracks: List[aiortc.MediaStreamTrack] = audio_producer_tracks + self.incoming_camera_tracks: Dict[str, Any] = {} + self.incoming_audio_tracks: List[Any] = [] + self.outgoing_video_tracks = video_producer_tracks + self.outgoing_audio_tracks = audio_producer_tracks self.should_add_data_channel = should_add_data_channel - self.messaging_channel: Optional[aiortc.RTCDataChannel] = None + self.messaging_channel: Optional[DataChannel] = None self.incoming_message_handlers: List[MessageHandler] = [] + self._consumer_tracks: List[Track] = [] + self._sender_tasks: List[asyncio.Task] = [] + self._track_state: List[Tuple[Track, TiciVideoStreamTrack, RtpPacketizationConfig]] = [] + self._receiver_reports: Dict[str, RtcpReceiverReport] = {} + self._receiver_report_tracks: Dict[str, Tuple[Track, int]] = {} self.incoming_media_ready_event = asyncio.Event() self.messaging_channel_ready_event = asyncio.Event() self.connection_attempted_event = asyncio.Event() self.connection_stopped_event = asyncio.Event() + self.gathering_complete_event = asyncio.Event() + self._loop: Optional[asyncio.AbstractEventLoop] = None - self.peer_connection.on("connectionstatechange", self._on_connectionstatechange) - self.peer_connection.on("datachannel", self._on_incoming_datachannel) - self.peer_connection.on("track", self._on_incoming_track) + self.peer_connection.on_state_change(self._on_connectionstatechange) + self.peer_connection.on_gathering_state_change(self._on_gatheringstatechange) + self.peer_connection.on_data_channel(self._on_incoming_datachannel) + if self.expected_incoming_camera_types or self.expected_incoming_audio: + self.peer_connection.on_track(self._on_incoming_track) self.logger = logging.getLogger("WebRTCStream") def _log_debug(self, msg: Any, *args): self.logger.debug(f"{type(self)}() {msg}", *args) + def _call_soon_threadsafe(self, fn: Callable, *args) -> None: + if self._loop is not None and self._loop.is_running(): + self._loop.call_soon_threadsafe(fn, *args) + else: + fn(*args) + + def _set_event(self, event: asyncio.Event) -> None: + self._call_soon_threadsafe(event.set) + @property def _number_of_incoming_media(self) -> int: media = len(self.incoming_camera_tracks) + len(self.incoming_audio_tracks) - # if stream does not add data_channel, then it means its incoming media += int(self.messaging_channel is not None) if not self.should_add_data_channel else 0 return media def _add_consumer_transceivers(self): - for _ in self.expected_incoming_camera_types: - self.peer_connection.addTransceiver("video", direction="recvonly") + for camera_type in self.expected_incoming_camera_types: + media = Description.Video(camera_type, Description.Direction.RecvOnly) + media.add_h264_codec(96) + track = self.peer_connection.add_track(media) + self._consumer_tracks.append(track) + self.incoming_camera_tracks[camera_type] = track if self.expected_incoming_audio: - self.peer_connection.addTransceiver("audio", direction="recvonly") + media = Description.Audio("audio", Description.Direction.RecvOnly) + media.add_opus_codec(111) + track = self.peer_connection.add_track(media) + self._consumer_tracks.append(track) + self.incoming_audio_tracks.append(track) - def _find_trackless_transceiver(self, kind: str) -> Optional[aiortc.RTCRtpTransceiver]: - transceivers = self.peer_connection.getTransceivers() - target_transceiver = None - for t in transceivers: - if t.kind == kind and t.sender.track is None: - target_transceiver = t - break + def _find_offer_video(self, remote_sdp: str, used_mids: set[str]) -> Tuple[str, int]: + desc = Description(remote_sdp, Description.Type.Offer) + for i in range(desc.media_count()): + media = desc.media(i) + if media is None or media.type() != "video" or media.mid() in used_mids: + continue + for payload_type in media.payload_types(): + with contextlib.suppress(ValueError): + rtp_map = media.rtp_map(payload_type) + if rtp_map is not None and rtp_map.format.upper() == "H264": + return media.mid(), payload_type + raise ValueError("Remote SDP does not offer H264 video") - return target_transceiver + def _make_video_media(self, track: TiciVideoStreamTrack, remote_sdp: str, used_mids: set[str]) -> Tuple[Description.Video, int, int, str]: + mid, payload_type = self._find_offer_video(remote_sdp, used_mids) + used_mids.add(mid) + ssrc = random.randint(1, 0xFFFFFFFF) + cname = f"teleoprtc-{random.getrandbits(32):08x}" + stream_id = f"stream-{random.getrandbits(32):08x}" + media = Description.Video(mid, Description.Direction.SendOnly) + media.add_h264_codec(payload_type) + media.add_ssrc(ssrc, cname, stream_id, track.id) + return media, ssrc, payload_type, cname - def _add_producer_tracks(self): + def _add_producer_tracks(self, remote_sdp: Optional[str] = None): + used_mids: set[str] = set() for track in self.outgoing_video_tracks: - target_transceiver = self._find_trackless_transceiver(track.kind) - if target_transceiver is None: - self.peer_connection.addTransceiver(track.kind, direction="sendonly") + media, ssrc, payload_type, cname = self._make_video_media(track, remote_sdp or "", used_mids) + rtc_track = self.peer_connection.add_track(media) - sender = self.peer_connection.addTrack(track) - if hasattr(track, "codec_preference") and track.codec_preference() is not None: - transceiver = next(t for t in self.peer_connection.getTransceivers() if t.sender == sender) - self._force_codec(transceiver, track.codec_preference(), "video") - for track in self.outgoing_audio_tracks: - target_transceiver = self._find_trackless_transceiver(track.kind) - if target_transceiver is None: - self.peer_connection.addTransceiver(track.kind, direction="sendonly") + rtp_config = RtpPacketizationConfig(ssrc, cname, payload_type, H264RtpPacketizer.CLOCK_RATE) + rtp_config.start_timestamp = random.randint(0, 0xFFFFFFFF) + rtp_config.timestamp = rtp_config.start_timestamp + rtp_config.sequence_number = random.randint(0, 0xFFFF) - self.peer_connection.addTrack(track) + packetizer = H264RtpPacketizer(NalUnit.Separator.LongStartSequence, rtp_config, 1200) + packetizer.add_to_chain(RtcpSrReporter(rtp_config)) + packetizer.add_to_chain(PliHandler(track.request_keyframe)) + packetizer.add_to_chain(RtcpNackResponder()) + rtc_track.set_media_handler(packetizer) - def _add_messaging_channel(self, channel: Optional[aiortc.RTCDataChannel] = None): - if not channel: - channel = self.peer_connection.createDataChannel("data", ordered=True) + camera_type, _ = parse_video_track_id(track.id) + rtc_track.reset_callbacks() + self._receiver_report_tracks[camera_type] = (rtc_track, ssrc) + self._track_state.append((rtc_track, track, rtp_config)) - for handler in self.incoming_message_handlers: - channel.on("message", handler) + if self.outgoing_audio_tracks: + raise NotImplementedError("Audio producer tracks are not implemented with libdatachannel") - if channel.readyState == "open": - self.messaging_channel_ready_event.set() - else: - channel.on("open", lambda: self.messaging_channel_ready_event.set()) + def _add_messaging_channel(self, channel: Optional[DataChannel] = None): + if channel is None: + channel = self.peer_connection.create_data_channel("data") self.messaging_channel = channel - def _force_codec(self, transceiver: aiortc.RTCRtpTransceiver, codec: str, stream_type: str): - codec_mime = f"{stream_type}/{codec.upper()}" - rtp_codecs = aiortc.RTCRtpSender.getCapabilities(stream_type).codecs - rtp_codec = [c for c in rtp_codecs if c.mimeType == codec_mime] - transceiver.setCodecPreferences(rtp_codec) + def on_message(message: Union[bytes, str]): + for handler in list(self.incoming_message_handlers): + self._call_soon_threadsafe(handler, message) - def _on_connectionstatechange(self): - self._log_debug("connection state is %s", self.peer_connection.connectionState) - if self.peer_connection.connectionState in ['connected', 'failed']: - self.connection_attempted_event.set() - if self.peer_connection.connectionState in ['disconnected', 'closed', 'failed']: - self.connection_stopped_event.set() + def on_open(): + self._set_event(self.messaging_channel_ready_event) - def _on_incoming_track(self, track: aiortc.MediaStreamTrack): - self._log_debug("got track: %s %s", track.kind, track.id) - if track.kind == "video": - camera_type, _ = parse_video_track_id(track.id) - if camera_type in self.expected_incoming_camera_types: - self.incoming_camera_tracks[camera_type] = track - elif track.kind == "audio": - if self.expected_incoming_audio: - self.incoming_audio_tracks.append(track) + def on_closed(): + self._set_event(self.connection_stopped_event) + + channel.on_message(on_message) + channel.on_open(on_open) + channel.on_closed(on_closed) + if channel.is_open(): + self._set_event(self.messaging_channel_ready_event) self._on_after_media() - def _on_incoming_datachannel(self, channel: aiortc.RTCDataChannel): - self._log_debug("got data channel: %s", channel.label) - if channel.label == "data" and self.messaging_channel is None: + def _retain_messaging_channel(self) -> None: + if self.messaging_channel is None: + return + + # No native callback can be running before a remote description is set. + if not self.messaging_channel_ready_event.is_set() and self.peer_connection.remote_description() is None: + self.messaging_channel = None + return + + if self._retain_messaging_channel_on_close: + self._retained_messaging_channels.append(self.messaging_channel) + self.messaging_channel = None + + def _on_connectionstatechange(self, state: PeerConnection.State): + self._log_debug("connection state is %s", state) + if state in (PeerConnection.State.Connected, PeerConnection.State.Failed): + self._set_event(self.connection_attempted_event) + if state in (PeerConnection.State.Disconnected, PeerConnection.State.Closed, PeerConnection.State.Failed): + self._set_event(self.connection_stopped_event) + + def _on_gatheringstatechange(self, state: PeerConnection.GatheringState): + self._log_debug("gathering state is %s", state) + if state == PeerConnection.GatheringState.Complete: + self._set_event(self.gathering_complete_event) + + def _on_incoming_track(self, track: Track): + self._log_debug("got track: %s", track.mid()) + try: + camera_type, _ = parse_video_track_id(track.mid()) + except ValueError: + camera_type = track.mid() + if camera_type in self.expected_incoming_camera_types: + self.incoming_camera_tracks[camera_type] = track + elif self.expected_incoming_audio: + self.incoming_audio_tracks.append(track) + self._on_after_media() + + def _on_incoming_datachannel(self, channel: DataChannel): + self._log_debug("got data channel: %s", channel.label()) + if channel.label() == "data" and self.messaging_channel is None: self._add_messaging_channel(channel) - self._on_after_media() + + def _update_receiver_report(self, camera_type: str, ssrc: int, message: bytes) -> None: + for report in _decode_receiver_reports(message): + if report.ssrc == ssrc: + self._receiver_reports[camera_type] = report def _on_after_media(self): - if self._number_of_incoming_media == self.expected_number_of_incoming_media: - self.incoming_media_ready_event.set() + if self.expected_number_of_incoming_media is not None and self._number_of_incoming_media >= self.expected_number_of_incoming_media: + self._set_event(self.incoming_media_ready_event) def _parse_incoming_streams(self, remote_sdp: str): - desc = aiortc.sdp.SessionDescription.parse(remote_sdp) - audio_video_media_count = len([m for m in desc.media if m.kind in ["audio", "video"] and m.direction in ["sendonly", "sendrecv"]]) - data_media_count = int(any(m for m in desc.media if m.kind == "application")) if not self.should_add_data_channel else 0 - self.expected_number_of_incoming_media = audio_video_media_count + data_media_count + desc = Description(remote_sdp, Description.Type.Offer) + media_count = 0 + for i in range(desc.media_count()): + media = desc.media(i) + if media is None: + continue + direction = media.direction() + if media.type() in ("audio", "video") and direction in (Description.Direction.SendOnly, Description.Direction.SendRecv): + media_count += 1 + data_media_count = int(desc.has_application()) if not self.should_add_data_channel else 0 + self.expected_number_of_incoming_media = media_count + data_media_count + if self.expected_number_of_incoming_media == 0: + self._set_event(self.incoming_media_ready_event) def has_incoming_video_track(self, camera_type: str) -> bool: return camera_type in self.incoming_camera_tracks @@ -158,65 +274,122 @@ class WebRTCBaseStream(abc.ABC): def has_messaging_channel(self) -> bool: return self.messaging_channel is not None - def get_incoming_video_track(self, camera_type: str, buffered: bool = False) -> aiortc.MediaStreamTrack: + def get_incoming_video_track(self, camera_type: str) -> Track: assert camera_type in self.incoming_camera_tracks, "Video tracks are not enabled on this stream" assert self.is_started, "Stream must be started" + return self.incoming_camera_tracks[camera_type] - track = self.incoming_camera_tracks[camera_type] - relay_track = self.media_relay.subscribe(track, buffered=buffered) - return relay_track - - def get_incoming_audio_track(self, buffered: bool = False) -> aiortc.MediaStreamTrack: + def get_incoming_audio_track(self) -> Track: assert len(self.incoming_audio_tracks) > 0, "Audio tracks are not enabled on this stream" assert self.is_started, "Stream must be started" + return self.incoming_audio_tracks[0] - track = self.incoming_audio_tracks[0] - relay_track = self.media_relay.subscribe(track, buffered=buffered) - return relay_track - - def get_messaging_channel(self) -> aiortc.RTCDataChannel: + def get_messaging_channel(self) -> DataChannel: assert self.messaging_channel is not None, "Messaging channel is not enabled on this stream" assert self.is_started, "Stream must be started" - return self.messaging_channel + def get_receiver_report_stats(self) -> Dict[str, RtcpReceiverReport]: + return dict(self._receiver_reports) + def set_message_handler(self, message_handler: MessageHandler): self.incoming_message_handlers.append(message_handler) - if self.messaging_channel is not None: - self.messaging_channel.on("message", message_handler) @property def is_started(self) -> bool: return self.peer_connection is not None and \ - self.peer_connection.localDescription is not None and \ - self.peer_connection.remoteDescription is not None and \ - self.peer_connection.connectionState != "closed" + self.peer_connection.local_description() is not None and \ + self.peer_connection.remote_description() is not None and \ + self.peer_connection.state() != PeerConnection.State.Closed @property def is_connected_and_ready(self) -> bool: return self.peer_connection is not None and \ - self.peer_connection.connectionState == "connected" and \ + self.peer_connection.state() == PeerConnection.State.Connected and \ (self.expected_number_of_incoming_media == 0 or self.incoming_media_ready_event.is_set()) + async def _wait_for_gathering_complete(self): + if self.peer_connection.gathering_state() != PeerConnection.GatheringState.Complete: + await self.gathering_complete_event.wait() + + async def _send_track_loop(self, rtc_track: Track, producer_track: TiciVideoStreamTrack, rtp_config: RtpPacketizationConfig): + while True: + if not rtc_track.is_open(): + await asyncio.sleep(0.01) + continue + + try: + packet = await producer_track.recv() + data = bytes(packet) + if not data: + continue + + pts = int(packet.pts or 0) + timestamp = (rtp_config.start_timestamp + pts) & 0xFFFFFFFF + rtc_track.send_frame(data, FrameInfo(timestamp)) + except asyncio.CancelledError: + raise + except Exception: + self.logger.exception("Error in send track loop for track %s", producer_track.id) + self._set_event(self.connection_stopped_event) + break + + async def _receiver_report_loop(self): + while True: + for camera_type, (rtc_track, ssrc) in self._receiver_report_tracks.items(): + for _ in range(32): + try: + message = rtc_track.receive() + if message is None: # go until queue empty (bounded to 32) + break + if isinstance(message, bytes): + self._update_receiver_report(camera_type, ssrc, message) + except asyncio.CancelledError: + raise + except Exception: + self.logger.exception("Error receiving report for %s", camera_type) + break + await asyncio.sleep(0.05) + + def _start_sender_tasks(self): + for rtc_track, producer_track, rtp_config in self._track_state: + self._sender_tasks.append(asyncio.create_task(self._send_track_loop(rtc_track, producer_track, rtp_config))) + if self._track_state: + self._sender_tasks.append(asyncio.create_task(self._receiver_report_loop())) + async def wait_for_connection(self): assert self.is_started await self.connection_attempted_event.wait() - if self.peer_connection.connectionState != 'connected': + if self.peer_connection.state() != PeerConnection.State.Connected: raise ValueError("Connection failed.") if self.expected_number_of_incoming_media: await self.incoming_media_ready_event.wait() if self.messaging_channel is not None: await self.messaging_channel_ready_event.wait() + self._start_sender_tasks() async def wait_for_disconnection(self): assert self.is_connected_and_ready, "Stream is not connected/ready yet (make sure wait_for_connection was awaited)" await self.connection_stopped_event.wait() async def stop(self): - await self.peer_connection.close() + for task in self._sender_tasks: + task.cancel() + for task in self._sender_tasks: + with contextlib.suppress(asyncio.CancelledError): + await task + self._sender_tasks.clear() + self._retain_messaging_channel() + self.peer_connection.close() + self.incoming_camera_tracks.clear() + self.incoming_audio_tracks.clear() + self._consumer_tracks.clear() + self._track_state.clear() + self._receiver_reports.clear() + self._receiver_report_tracks.clear() @abc.abstractmethod - async def start(self) -> aiortc.RTCSessionDescription: + async def start(self) -> RTCSessionDescription: raise NotImplementedError @@ -225,76 +398,46 @@ class WebRTCOfferStream(WebRTCBaseStream): super().__init__(*args, **kwargs) self.session_provider = session_provider - async def start(self) -> aiortc.RTCSessionDescription: + async def start(self) -> RTCSessionDescription: + self._loop = asyncio.get_running_loop() self._add_consumer_transceivers() if self.should_add_data_channel: self._add_messaging_channel() - self._add_producer_tracks() - offer = await self.peer_connection.createOffer() - await self.peer_connection.setLocalDescription(offer) - actual_offer = self.peer_connection.localDescription + self.peer_connection.set_local_description(Description.Type.Offer) + await self._wait_for_gathering_complete() + actual_offer = self.peer_connection.local_description() streaming_offer = StreamingOffer( - sdp=actual_offer.sdp, + sdp=str(actual_offer), video=list(self.expected_incoming_camera_types), ) remote_answer = await self.session_provider(streaming_offer) self._parse_incoming_streams(remote_sdp=remote_answer.sdp) - await self.peer_connection.setRemoteDescription(remote_answer) - actual_answer = self.peer_connection.remoteDescription + self.peer_connection.set_remote_description(Description(remote_answer.sdp, Description.Type.Answer)) + self._on_after_media() + actual_answer = self.peer_connection.remote_description() - return actual_answer + return RTCSessionDescription(str(actual_answer), actual_answer.type_string()) class WebRTCAnswerStream(WebRTCBaseStream): - def __init__(self, session: aiortc.RTCSessionDescription, *args, **kwargs): + _retain_messaging_channel_on_close = True + + def __init__(self, session: RTCSessionDescription, *args, **kwargs): super().__init__(*args, **kwargs) self.session = session - def _probe_video_codecs(self) -> List[str]: - codecs = [] - for track in self.outgoing_video_tracks: - if hasattr(track, "codec_preference") and track.codec_preference() is not None: - codecs.append(track.codec_preference()) - - return codecs - - def _override_incoming_video_codecs(self, remote_sdp: str, codecs: List[str]) -> str: - desc = aiortc.sdp.SessionDescription.parse(remote_sdp) - codec_mimes = [f"video/{c}" for c in codecs] - for m in desc.media: - if m.kind != "video": - continue - - preferred_codecs: List[aiortc.RTCRtpCodecParameters] = [c for c in m.rtp.codecs if c.mimeType in codec_mimes] - if len(preferred_codecs) == 0: - raise ValueError(f"None of {preferred_codecs} codecs is supported in remote SDP") - - m.rtp.codecs = preferred_codecs - m.fmt = [c.payloadType for c in preferred_codecs] - - return str(desc) - - async def start(self) -> aiortc.RTCSessionDescription: - assert self.peer_connection.remoteDescription is None, "Connection already established" - - self._add_consumer_transceivers() - - # since we sent already encoded frames in some cases (e.g. livestream video tracks are in H264), we need to force aiortc to actually use it - # we do that by overriding supported codec information on incoming sdp - preferred_codecs = self._probe_video_codecs() - if len(preferred_codecs) > 0: - self.session.sdp = self._override_incoming_video_codecs(self.session.sdp, preferred_codecs) + async def start(self) -> RTCSessionDescription: + self._loop = asyncio.get_running_loop() + assert self.peer_connection.remote_description() is None, "Connection already established" self._parse_incoming_streams(remote_sdp=self.session.sdp) - await self.peer_connection.setRemoteDescription(self.session) + self.peer_connection.set_remote_description(Description(self.session.sdp, Description.Type.Offer)) + self._add_producer_tracks(self.session.sdp) - self._add_producer_tracks() - - answer = await self.peer_connection.createAnswer() - await self.peer_connection.setLocalDescription(answer) - actual_answer = self.peer_connection.localDescription - - return actual_answer + self.peer_connection.set_local_description(Description.Type.Answer) + await self._wait_for_gathering_complete() + actual_answer = self.peer_connection.local_description() + return RTCSessionDescription(str(actual_answer), actual_answer.type_string()) diff --git a/teleoprtc_repo/teleoprtc/tracks.py b/teleoprtc_repo/teleoprtc/tracks.py index 10b234aa2..12bcc88e7 100644 --- a/teleoprtc_repo/teleoprtc/tracks.py +++ b/teleoprtc_repo/teleoprtc/tracks.py @@ -1,11 +1,11 @@ -import asyncio -import logging -import time import fractions -from typing import Any, Optional, Tuple +import logging +import uuid +from typing import Any, Tuple -import aiortc -from aiortc.mediastreams import VIDEO_CLOCK_RATE, VIDEO_TIME_BASE + +VIDEO_CLOCK_RATE = 90000 +VIDEO_TIME_BASE = fractions.Fraction(1, VIDEO_CLOCK_RATE) def video_track_id(camera_type: str, track_id: str) -> str: @@ -21,57 +21,51 @@ def parse_video_track_id(track_id: str) -> Tuple[str, str]: return camera_type, track_id -class TiciVideoStreamTrack(aiortc.MediaStreamTrack): +class TiciVideoStreamTrack: """ - Abstract video track which associates video track with camera_type + Abstract video track which associates video track with camera_type. """ kind = "video" def __init__(self, camera_type: str, dt: float, time_base: fractions.Fraction = VIDEO_TIME_BASE, clock_rate: int = VIDEO_CLOCK_RATE): assert camera_type in ["driver", "wideRoad", "road"] - super().__init__() - # override track id to include camera type - client needs that for identification - self._id: str = video_track_id(camera_type, self._id) - self._dt: float = dt + self._id: str = video_track_id(camera_type, str(uuid.uuid4())) self._time_base: fractions.Fraction = time_base self._clock_rate: int = clock_rate - self._start: Optional[float] = None self._logger = logging.getLogger("WebRTCStream") + self.readyState = "live" + + @property + def id(self) -> str: + return self._id + + def stop(self) -> None: + self.readyState = "ended" def log_debug(self, msg: Any, *args): self._logger.debug(f"{type(self)}() {msg}", *args) - async def next_pts(self, current_pts) -> float: - pts: float = current_pts + self._dt * self._clock_rate + async def recv(self): + raise NotImplementedError() - data_time = pts * self._time_base - if self._start is None: - self._start = time.time() - data_time - else: - wait_time = self._start + data_time - time.time() - await asyncio.sleep(wait_time) - - return pts - - def codec_preference(self) -> Optional[str]: - return None + def request_keyframe(self) -> None: + pass -class TiciTrackWrapper(aiortc.MediaStreamTrack): +class TiciTrackWrapper(TiciVideoStreamTrack): """ - Associates video track with camera_type + Associates a generic video track with camera_type. """ - def __init__(self, camera_type: str, track: aiortc.MediaStreamTrack): + def __init__(self, camera_type: str, track: Any): assert track.kind == "video" - assert not isinstance(track, TiciVideoStreamTrack) - super().__init__() + super().__init__(camera_type, getattr(track, "_dt", 0.05)) self._id = video_track_id(camera_type, track.id) self._track = track - @property - def kind(self) -> str: - return self._track.kind - async def recv(self): return await self._track.recv() + def stop(self) -> None: + super().stop() + if hasattr(self._track, "stop"): + self._track.stop() diff --git a/tools/agnos/flash_desktop_system_to_comma.sh b/tools/agnos/flash_desktop_system_to_comma.sh index e13580b5d..7d91188ad 100755 --- a/tools/agnos/flash_desktop_system_to_comma.sh +++ b/tools/agnos/flash_desktop_system_to_comma.sh @@ -2,12 +2,45 @@ set -euo pipefail HOST="${1:-comma@192.168.3.110}" -IMAGE="${2:-/Users/dominickthompson/Desktop/system8.img.xz}" +IMAGE="${2:-/Users/dominickthompson/Desktop/system17.img.xz}" +METADATA="${3:-${IMAGE}.metadata.json}" SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)" REPO_ROOT="$(cd -- "${SCRIPT_DIR}/../.." && pwd)" -SSH_KEY="${SSH_KEY:-${REPO_ROOT}/system/hardware/tici/id_rsa}" -SSH_OPTS=(-i "$SSH_KEY" -o BatchMode=yes -o StrictHostKeyChecking=no -o UserKnownHostsFile=/dev/null) +SSH_OPTS=(-o BatchMode=yes -o StrictHostKeyChecking=no -o UserKnownHostsFile=/dev/null) +if [[ -n "${SSH_KEY:-}" ]]; then + SSH_OPTS=(-i "$SSH_KEY" -o IdentitiesOnly=yes "${SSH_OPTS[@]}") +fi + +for required_path in "$IMAGE" "$METADATA"; do + [[ -f "$required_path" ]] || { echo "missing file: $required_path" >&2; exit 1; } +done + +metadata_value() { + python3 - "$METADATA" "$1" <<'PY' +import json +import sys +from pathlib import Path + +metadata = json.loads(Path(sys.argv[1]).read_text(encoding="utf-8")) +value = metadata[sys.argv[2]] +if isinstance(value, (dict, list)): + raise SystemExit(f"metadata field {sys.argv[2]} is not scalar") +print(value) +PY +} + +BASE_VERSION="$(metadata_value base_version)" +EXPECTED_VERSION="$(metadata_value target_version)" +RAW_HASH="$(metadata_value raw_sha256)" +RAW_SIZE="$(metadata_value raw_size)" +EXPECTED_XZ_HASH="$(metadata_value xz_sha256)" +ACTUAL_XZ_HASH="$(shasum -a 256 "$IMAGE" | awk '{print $1}')" + +[[ "$ACTUAL_XZ_HASH" == "$EXPECTED_XZ_HASH" ]] || { + echo "compressed image hash mismatch: got $ACTUAL_XZ_HASH, expected $EXPECTED_XZ_HASH" >&2 + exit 1 +} SESSION="local_agnos_flash" REMOTE_DIR="/data/local_agnos_flash" @@ -15,40 +48,49 @@ REMOTE_MANIFEST="${REMOTE_DIR}/agnos-local-system.json" REMOTE_RUNNER="${REMOTE_DIR}/run_flash.sh" REMOTE_AGNOS="${REMOTE_DIR}/agnos.py" PORT="8989" - -EXPECTED_VERSION="12.8.28" -RAW_HASH="4c01245932068aedfceb41cb1aab1f7f044f6659aa2fe2de558f99e2d3aa5793" -RAW_SIZE="5368709120" - -if [[ ! -f "$IMAGE" ]]; then - echo "missing image: $IMAGE" >&2 - exit 1 -fi - IMAGE_NAME="$(basename "$IMAGE")" REMOTE_IMAGE="${REMOTE_DIR}/${IMAGE_NAME}" +INSTALLED_VERSION="$(ssh "${SSH_OPTS[@]}" "$HOST" 'tr -d "\r\n" &2 + exit 1 + ;; +esac + +echo "[CHECK] Device AGNOS: $INSTALLED_VERSION" +echo "[CHECK] Candidate AGNOS: $EXPECTED_VERSION" +echo "[CHECK] Candidate XZ hash: $ACTUAL_XZ_HASH" + ssh "${SSH_OPTS[@]}" "$HOST" "mkdir -p '$REMOTE_DIR'" scp "${SSH_OPTS[@]}" "$IMAGE" "$HOST:$REMOTE_IMAGE" LOCAL_AGNOS="$(mktemp "${TMPDIR:-/tmp}/agnos-local.XXXXXX.py")" +trap 'rm -f "$LOCAL_AGNOS"' EXIT python3 - "$REPO_ROOT/system/hardware/tici/agnos.py" "$LOCAL_AGNOS" <<'PY' import sys from pathlib import Path src, dst = map(Path, sys.argv[1:]) data = src.read_text(encoding="utf-8") +needle = "import openpilot.system.updated.casync.casync as casync" +if data.count(needle) != 1: + raise SystemExit("could not isolate the unused casync dependency in agnos.py") data = data.replace( - "import openpilot.system.updated.casync.casync as casync", - """class _UnusedCasync: + needle, + '''class _UnusedCasync: + ChunkReader = object + ChunkDict = object + def __getattr__(self, name): raise RuntimeError("casync support is unavailable in local AGNOS flash runner") -casync = _UnusedCasync()""", +casync = _UnusedCasync()''', ) dst.write_text(data, encoding="utf-8") PY scp "${SSH_OPTS[@]}" "$LOCAL_AGNOS" "$HOST:$REMOTE_AGNOS" -rm -f "$LOCAL_AGNOS" ssh "${SSH_OPTS[@]}" "$HOST" "cat > '$REMOTE_MANIFEST'" < >(tee -a "${REMOTE_DIR}/flash.log") 2>&1 -echo "[STEP] Local AGNOS system flash" echo "[CHECK] Installed AGNOS: $(cat /VERSION 2>/dev/null || echo unknown)" echo "[CHECK] Target AGNOS: ${EXPECTED_VERSION}" +echo "[CHECK] Active slot: $(abctl --boot_slot)" +df -h /data -if [[ ! -f "$REMOTE_AGNOS" ]]; then - echo "[ERROR] $REMOTE_AGNOS not found" >&2 - exit 1 -fi - -if [[ -x /usr/local/venv/bin/python3 ]]; then - PYTHON_BIN="/usr/local/venv/bin/python3" -else - PYTHON_BIN="python3" -fi +PYTHON_BIN="/usr/local/venv/bin/python3" +[[ -x "$PYTHON_BIN" ]] || { echo "[ERROR] managed Python is unavailable" >&2; exit 1; } pkill -f "http.server ${PORT}.*${REMOTE_DIR}" >/dev/null 2>&1 || true "$PYTHON_BIN" -m http.server "$PORT" --bind 127.0.0.1 --directory "$REMOTE_DIR" >"${REMOTE_DIR}/http.log" 2>&1 & http_pid="$!" trap 'kill "$http_pid" >/dev/null 2>&1 || true' EXIT -http_ready=0 for _ in $(seq 1 20); do - if "$PYTHON_BIN" - "${PORT}" "${IMAGE_NAME}" <<'PY' + if "$PYTHON_BIN" - "$REMOTE_MANIFEST" <<'PY' +import json import sys import urllib.request +from pathlib import Path -port, image_name = sys.argv[1], sys.argv[2] -with urllib.request.urlopen(f"http://127.0.0.1:{port}/{image_name}", timeout=2) as resp: - resp.read(1) +for entry in json.loads(Path(sys.argv[1]).read_text(encoding="utf-8")): + with urllib.request.urlopen(entry["url"], timeout=2) as response: + response.read(1) PY then http_ready=1 @@ -115,24 +150,23 @@ PY sleep 0.25 done -if [[ "$http_ready" != "1" ]]; then - echo "[ERROR] Local image HTTP server did not become ready" >&2 +[[ "${http_ready:-0}" == "1" ]] || { + echo "[ERROR] local image server did not become ready" >&2 cat "${REMOTE_DIR}/http.log" >&2 || true exit 1 -fi +} -echo "[FLASH] Flashing local system image to inactive AGNOS slot" +echo "[FLASH] Writing and verifying the candidate in the inactive system slot" PYTHONPATH="$(dirname "$REMOTE_AGNOS")" "$PYTHON_BIN" "$REMOTE_AGNOS" --swap "$REMOTE_MANIFEST" -echo "[DONE] AGNOS flashed and slot swapped" -echo "[REBOOT] Rebooting now" +echo "[DONE] Candidate written, verified, and selected" sudo reboot REMOTE_RUNNER ssh "${SSH_OPTS[@]}" "$HOST" "tmux kill-session -t '$SESSION' >/dev/null 2>&1 || true" ssh "${SSH_OPTS[@]}" "$HOST" "rm -f '$REMOTE_DIR/flash.log' '$REMOTE_DIR/http.log'" ssh "${SSH_OPTS[@]}" "$HOST" \ - "tmux new-session -d -s '$SESSION' \"REMOTE_DIR='$REMOTE_DIR' REMOTE_MANIFEST='$REMOTE_MANIFEST' REMOTE_AGNOS='$REMOTE_AGNOS' PORT='$PORT' IMAGE_NAME='$IMAGE_NAME' EXPECTED_VERSION='$EXPECTED_VERSION' bash '$REMOTE_RUNNER'\"" + "tmux new-session -d -s '$SESSION' \"REMOTE_DIR='$REMOTE_DIR' REMOTE_MANIFEST='$REMOTE_MANIFEST' REMOTE_AGNOS='$REMOTE_AGNOS' PORT='$PORT' EXPECTED_VERSION='$EXPECTED_VERSION' bash '$REMOTE_RUNNER'\"" echo "Started remote tmux session: $SESSION" -echo "Watch it with: ssh $HOST 'tmux attach -t $SESSION'" +echo "After reboot, run tools/agnos/validate_agnos_runtime.sh $EXPECTED_VERSION on the device." diff --git a/tools/agnos/patch_system_reset_image.py b/tools/agnos/patch_system_reset_image.py old mode 100644 new mode 100755 index 9b9f1c595..1bd27251c --- a/tools/agnos/patch_system_reset_image.py +++ b/tools/agnos/patch_system_reset_image.py @@ -1,50 +1,149 @@ #!/usr/bin/env python3 +"""Build StarPilot AGNOS from the exact upstream system image. + +The output starts with comma's pinned AGNOS system partition, adds the Python +packages required by StarPilot's older runtime and C3 support, and customizes +only the stock setup/installer pair needed for StarPilot's factory install. +Reset, updater, Magic, NetworkManager, and every existing upstream Python +package remain byte-identical to the pinned image. +""" + import argparse import hashlib import json +import lzma import os import re import shutil -import subprocess import struct +import subprocess import tempfile import urllib.request import zipfile -from io import BytesIO from pathlib import Path -RESET_PATH_IN_IMAGE = "/usr/comma/reset" -COMMA_SH_PATH_IN_IMAGE = "/usr/comma/comma.sh" -MAGIC_PATH_IN_IMAGE = "/usr/comma/magic.py" -SETUP_PATH_IN_IMAGE = "/usr/comma/setup" -UPDATER_PATH_IN_IMAGE = "/usr/comma/updater" -BG_PATH_IN_IMAGE = "/usr/comma/bg.jpg" -WESTON_SERVICE_PATH_IN_IMAGE = "/lib/systemd/system/weston.service" -RESET_ENTRY_IN_ZIPAPP = "openpilot/system/ui/reset.py" -MICI_RESET_ENTRY_IN_ZIPAPP = "openpilot/system/ui/mici_reset.py" -TICI_RESET_ENTRY_IN_ZIPAPP = "openpilot/system/ui/tici_reset.py" -APPLICATION_ENTRY_IN_ZIPAPP = "openpilot/system/ui/lib/application.py" -WIFI_MANAGER_ENTRY_IN_SETUP_ZIPAPP = "openpilot/system/ui/lib/wifi_manager.py" -SETUP_ENTRY_IN_SETUP_ZIPAPP = "openpilot/system/ui/setup.py" -TICI_SETUP_ENTRY_IN_SETUP_ZIPAPP = "openpilot/system/ui/tici_setup.py" -MICI_SETUP_ENTRY_IN_SETUP_ZIPAPP = "openpilot/system/ui/mici_setup.py" -UPDATER_ENTRY_IN_ZIPAPP = "openpilot/system/ui/updater.py" VERSION_PATH_IN_IMAGE = "/VERSION" -PYTHON_SITE_PACKAGES_PATH_IN_IMAGE = "/usr/local/venv/lib/python3.12/site-packages" -AMDGPU_FIRMWARE_PATH_IN_IMAGE = "/lib/firmware/amdgpu" -PATCH_MARKER = "STARPILOT_C4_RESET_LAYOUT_V1" -APP_PATCH_MARKER = "STARPILOT_C4_RESET_APP_DIMENSIONS_V1" -SETUP_WIFI_PATCH_MARKER = "JEEPNY_AVAILABLE = True" -SETUP_BRANDING_PATCH_MARKER = "STARPILOT_SETUP_BRANDING_V1" -SETUP_SSH_RESTORE_PATCH_MARKER = "STARPILOT_SETUP_SSH_RESTORE_V1" -WESTON_BG_PATCH_MARKER = "STARPILOT_WESTON_BG_ORIENTATION_V2" -COMMA_SH_DISPLAY_WAIT_PATCH_MARKER = "STARPILOT_DISPLAY_READY_WAIT_V1" -JEEPNY_VERSION = "0.9.0" -JEEPNY_WHEEL_URL = "https://files.pythonhosted.org/packages/b2/a3/e137168c9c44d18eff0376253da9f1e9234d0239e0ee230d2fee6cea8e55/jeepney-0.9.0-py3-none-any.whl" -JEEPNY_WHEEL_SHA256 = "97e5714520c16fc0a45695e5365a2e11b81ea79bba796e26f9f1d178cb182683" -JEEPNY_PACKAGE_DIR = "jeepney" -JEEPNY_DIST_INFO_DIR = f"jeepney-{JEEPNY_VERSION}.dist-info" +SITE_PACKAGES_PATH_IN_IMAGE = "/usr/local/venv/lib/python3.12/site-packages" +LEGACY_RUNTIME_LIBRARY_DIR = "/usr/local/lib" +SETUP_PATH_IN_IMAGE = "/usr/comma/setup" +INSTALLER_PATH_IN_IMAGE = "/usr/comma/installer" +STAR_PILOT_GIT_URL = "https://github.com/firestar5683/openpilot.git" +STAR_PILOT_BRANCH = "StarPilot" +STAR_PILOT_DEPENDENCY_NAMES = ( + # C3/runtime compatibility + "crcmod", "crcmod-1.7.dist-info", "serial", "pyserial-3.5.dist-info", + "kaitaistruct.py", "kaitaistruct-0.11.dist-info", + # StarPilot always-on/default features + "cv2", "opencv_python_headless-4.11.0.86.dist-info", "opencv_python_headless.libs", + "mapbox_earcut.cpython-312-aarch64-linux-gnu.so", "mapbox_earcut-1.0.3.dist-info", + "jsonrpc", "json_rpc-1.15.0.dist-info", "xattr", "xattr-1.2.0.dist-info", + "onnx", "onnx-1.18.0.dist-info", "google", "protobuf-7.35.1.dist-info", + "typing_extensions.py", "typing_extensions-4.16.0.dist-info", + # Existing body/web tools still use aiohttp and PyAudio. The aiortc stack is + # deliberately not copied; WebRTC uses upstream's libdatachannel backend. + "aiohappyeyeballs", "aiohappyeyeballs-2.7.1.dist-info", + "aiohttp", "aiohttp-3.12.15.dist-info", + "aiosignal", "aiosignal-1.4.0.dist-info", + "attr", "attrs", "attrs-26.1.0.dist-info", + "frozenlist", "frozenlist-1.8.0.dist-info", + "multidict", "multidict-6.7.1.dist-info", + "propcache", "propcache-0.5.2.dist-info", + "yarl", "yarl-1.24.5.dist-info", + "pyaudio", "pyaudio-0.2.14.dist-info", +) +STAR_PILOT_DEPENDENCY_PATHS = tuple( + f"{SITE_PACKAGES_PATH_IN_IMAGE}/{name}" for name in STAR_PILOT_DEPENDENCY_NAMES +) +C3_DEPENDENCY_PATHS = tuple( + f"{SITE_PACKAGES_PATH_IN_IMAGE}/{name}" + for name in ("crcmod", "crcmod-1.7.dist-info", "serial", "pyserial-3.5.dist-info", "kaitaistruct.py", "kaitaistruct-0.11.dist-info") +) +LEGACY_RUNTIME_LIBRARY_NAMES = ( + # Existing StarPilot prebuilts use these legacy SONAMEs. Keep only their + # runtime closure; current source builds use upstream's managed packages. + "libcapnp-1.0.2.so", "libkj-1.0.2.so", + "libavformat.so.58", "libavformat.so.58.29.100", + "libavcodec.so.58", "libavcodec.so.58.54.100", + "libavutil.so.56", "libavutil.so.56.31.100", + "libswresample.so.3", "libswresample.so.3.5.100", +) +LEGACY_RUNTIME_LIBRARY_PATHS = tuple( + f"{LEGACY_RUNTIME_LIBRARY_DIR}/{name}" for name in LEGACY_RUNTIME_LIBRARY_NAMES +) +FACTORY_INSTALL_PATHS = frozenset({SETUP_PATH_IN_IMAGE, INSTALLER_PATH_IN_IMAGE}) +ALLOWED_IMAGE_MUTATIONS = frozenset({ + VERSION_PATH_IN_IMAGE, + *STAR_PILOT_DEPENDENCY_PATHS, + *LEGACY_RUNTIME_LIBRARY_PATHS, + *FACTORY_INSTALL_PATHS, +}) + +# Exact system partition pinned by ~/openpilot as of the 19.6 AGNOS release. +UPSTREAM_VERSION = "19.6" +UPSTREAM_SYSTEM_URL = ( + "https://commadist.azureedge.net/agnosupdate/" + "system-5b6ce7965904a157fd3a134ccfcb854f9ca5c1cc2a26b7cb80a4fa4e1cc4aaa3.img.xz" +) +UPSTREAM_RAW_SHA256 = "5b6ce7965904a157fd3a134ccfcb854f9ca5c1cc2a26b7cb80a4fa4e1cc4aaa3" +UPSTREAM_RAW_SIZE = 4_718_592_000 +UPSTREAM_SITE_PACKAGES_COUNT = 213 + +# Compatibility packages are copied from StarPilot's exact, previously +# deployed and field-tested 19.6.2 image. Existing upstream paths are never +# overwritten. +C3_DEPENDENCY_SOURCE_URL = ( + "https://www.dropbox.com/scl/fi/pewhzpqzi3aewuiaffc6m/system10.img.xz" + "?rlkey=olzrzulhs93zzghnjrskmdwxt&st=exnfk2oz&dl=1" +) +C3_DEPENDENCY_SOURCE_RAW_SHA256 = "ab395d4c963a908ab86709f1a6580a62dd24cb34ee71cf5fd5ec29d7d48d0e10" +C3_DEPENDENCY_SOURCE_RAW_SIZE = 4_718_592_000 +CANDIDATE_SITE_PACKAGES_COUNT = UPSTREAM_SITE_PACKAGES_COUNT + len(STAR_PILOT_DEPENDENCY_PATHS) + +# Hashes from that exact upstream image. Equality keeps all recovery plumbing +# stock except the explicitly customized setup/installer pair. +PROTECTED_PAYLOAD_HASHES = { + "/etc/NetworkManager/NetworkManager.conf": "779db62d2d4c5f8ce504c5d1f2994d34a9f35296d5efb7f3a48cb1e8a0d4778e", + "/etc/NetworkManager/conf.d/10-globally-managed-devices.conf": "45e653e2f709c027fad41f2d86b70e008b72c6bf4d34590b4765ebe8fe3ea948", + "/lib/systemd/system/NetworkManager.service": "fb33a80bf8c78b3af004d4b294c47a0139e37742c1d0d5a6a7663c7d1f4a2b48", + "/usr/comma/updater": "9df4edbeb5849de03f9c2d691d04646af84a3ef74c2f33be8e73d9281daebe99", + "/usr/comma/reset": "97ed6413515d0674442c42ae6e20baccf66dd6bb4ec382ee4cf0cc5ebe84e739", + "/usr/comma/magic.py": "c4416e66b127b31c17d08e6723ad46d12af7683e56626512d66c708b5f347ac9", + "/usr/comma/setup_keys": "934f74ab4b2ac06048418c2857be3a041e192ec03c09979987691c23c91353bd", + "/usr/comma/comma.sh": "bcba2b336cf0ca852786f8a58bbce407e0e9fe952c26fc5d903f6d9a34b44b4f", +} +UPSTREAM_FACTORY_INSTALL_HASHES = { + INSTALLER_PATH_IN_IMAGE: "85f6d9e54286a3842920d6967b187478b4e43d6171c331d72d3fb3102106e101", + SETUP_PATH_IN_IMAGE: "c382ce266653bad781c25e403ddea4af508aa6f3ea2eef3f568d964982fad9d6", +} +UPSTREAM_PAYLOAD_HASHES = {**PROTECTED_PAYLOAD_HASHES, **UPSTREAM_FACTORY_INSTALL_HASHES} + +SETUP_SOURCE_MEMBERS = ( + "openpilot/system/ui/mici_setup.py", + "openpilot/system/ui/tici_setup.py", +) + +UPSTREAM_REQUIRED_VENV_PATHS = { + "capnp": f"{SITE_PACKAGES_PATH_IN_IMAGE}/capnp", + "numpy": f"{SITE_PACKAGES_PATH_IN_IMAGE}/numpy", + "Crypto": f"{SITE_PACKAGES_PATH_IN_IMAGE}/Crypto", + "tqdm": f"{SITE_PACKAGES_PATH_IN_IMAGE}/tqdm", + "raylib": f"{SITE_PACKAGES_PATH_IN_IMAGE}/raylib", +} +REQUIRED_VENV_PATHS = { + **UPSTREAM_REQUIRED_VENV_PATHS, + "crcmod": f"{SITE_PACKAGES_PATH_IN_IMAGE}/crcmod", + "serial": f"{SITE_PACKAGES_PATH_IN_IMAGE}/serial", + "kaitaistruct": f"{SITE_PACKAGES_PATH_IN_IMAGE}/kaitaistruct.py", + "cv2": f"{SITE_PACKAGES_PATH_IN_IMAGE}/cv2", + "mapbox_earcut": f"{SITE_PACKAGES_PATH_IN_IMAGE}/mapbox_earcut.cpython-312-aarch64-linux-gnu.so", + "jsonrpc": f"{SITE_PACKAGES_PATH_IN_IMAGE}/jsonrpc", + "xattr": f"{SITE_PACKAGES_PATH_IN_IMAGE}/xattr", + "onnx": f"{SITE_PACKAGES_PATH_IN_IMAGE}/onnx", + "aiohttp": f"{SITE_PACKAGES_PATH_IN_IMAGE}/aiohttp", + "pyaudio": f"{SITE_PACKAGES_PATH_IN_IMAGE}/pyaudio", +} + ANDROID_SPARSE_MAGIC = 0xED26FF3A CHUNK_TYPE_RAW = 0xCAC1 CHUNK_TYPE_FILL = 0xCAC2 @@ -52,1312 +151,878 @@ CHUNK_TYPE_DONT_CARE = 0xCAC3 CHUNK_TYPE_CRC32 = 0xCAC4 XZ_MAGIC = b"\xFD7zXZ\x00" -AMDGPU_FIRMWARE_SHA256 = { - "gc_12_0_0_imu.bin.zst": "aa15e5b3156bffc45e0c50bccbcd364fbd3f958531b695b7487a803d780b8328", - "gc_12_0_0_me.bin.zst": "d7eba5197f2580f32b8256b1d9cb68e723e9e644293a34446a7913e3c093cba5", - "gc_12_0_0_mec.bin.zst": "1931593440b8f9423580d9e2cdc5b34e7c682cdffe1ca4b74b0c2f6a0420236d", - "gc_12_0_0_pfp.bin.zst": "16bfd64c10fe73b5e760055069a60e5841dba16c0ed4edb56c20d675e23901f6", - "gc_12_0_0_rlc.bin.zst": "6436b582734a413456fff3d3c7195e71cc9e78a7ed31ee21c83ffd6fae1ad186", - "psp_14_0_2_sos.bin.zst": "7b538448b57d4f9dd06b2eea90d4f86a16e65e3027cdecee8db71c2c5f1fa243", - "sdma_7_0_0.bin.zst": "beaafb53993a106edd392392d5896245ae2a957c6d0f495d0002eec72ad8ad38", - "smu_14_0_2.bin.zst": "6951995d1d606f4dc60c895f19d34ed18aa40e62129f83d8510c45e8aa9ae2fc", -} - -DEFAULT_SYNC_COMMA_FILES = [ - "/usr/comma/bg.jpg", - "/usr/comma/comma.sh", - "/usr/comma/debug.py", - "/usr/comma/fs_setup.sh", - "/usr/comma/installer", - "/usr/comma/magic.py", - "/usr/comma/power_drop_monitor.py", - "/usr/comma/power_monitor.py", - "/usr/comma/reset", - "/usr/comma/screen_calibration.py", - "/usr/comma/setup", - "/usr/comma/setup_keys", - "/usr/comma/updater", -] - -INODE_MODE_TYPE_PREFIX = { - "regular": "100", - "directory": "040", - "symlink": "120", - "character": "20", - "block": "60", - "fifo": "10", - "socket": "140", -} - -def patch_application_script(original: bytes) -> bytes: - if APP_PATCH_MARKER.encode("utf-8") in original: - return original - - text = original.decode("utf-8", "replace") - replacement = ( - f"# {APP_PATCH_MARKER}\n" - "_dt = HARDWARE.get_device_type()\n" - "if _dt in ('tici', 'tizi'):\n" - " gui_app = GuiApplication(2160, 1080)\n" - "else:\n" - " gui_app = GuiApplication(536, 240)\n" - ) - - fixed = text.replace("gui_app = GuiApplication(2160, 1080)", replacement) - if fixed == text: - fixed = re.sub( - r"^gui_app\s*=\s*GuiApplication\([^\n]+\)\s*$", - replacement.rstrip("\n"), - text, - count=1, - flags=re.MULTILINE, - ) - - if fixed == text: - # Newer upstream application.py is already device-aware via GuiApplication defaults. - # Keep it as-is and stamp a marker so verification can still pass. - if text.startswith("#!"): - first_nl = text.find("\n") - if first_nl != -1: - fixed = text[:first_nl + 1] + f"# {APP_PATCH_MARKER} (no-op)\n" + text[first_nl + 1:] - else: - fixed = text + f"\n# {APP_PATCH_MARKER} (no-op)\n" - else: - fixed = f"# {APP_PATCH_MARKER} (no-op)\n" + text - - return fixed.encode("utf-8") - - -def patch_setup_wifi_manager() -> bytes: - """ - Replace setup zipapp's wifi_manager.py with repo version that gracefully handles - missing jeepney (fallback to nmcli/fake), preventing setup boot-logo hangs. - """ - repo_root = Path(__file__).resolve().parents[2] - src = repo_root / "system/ui/lib/wifi_manager.py" - if not src.is_file(): - raise RuntimeError(f"Unable to find repo wifi_manager source: {src}") - data = src.read_bytes() - if SETUP_WIFI_PATCH_MARKER.encode("utf-8") not in data: - raise RuntimeError("Repo wifi_manager.py does not appear to include jeepney fallback marker") - return data - - -def patched_weston_bg_python() -> str: - som_id_path = "/sys/devices/platform/vendor/vendor:gpio-som-id/som_id" - return ( - "from PIL import Image; " - f"som=open(\"{som_id_path}\").read().strip(); " - "img=Image.open(\"/usr/comma/bg.jpg\").convert(\"RGB\"); " - "img=img.rotate(180) if som == \"1\" else img; " - "mask=img.convert(\"L\").point(lambda p: 255 if p > 16 else 0); " - "bbox=mask.getbbox(); " - "logo=img.crop(bbox) if bbox else img; " - # Pillow positive degrees are counter-clockwise; -90 pre-rotates the source clockwise. - "logo=logo.rotate(-90, expand=True); " - "resample=Image.Resampling.LANCZOS if hasattr(Image, \"Resampling\") else Image.LANCZOS; " - "logo=logo.resize((max(1, logo.width//3), max(1, logo.height//3)), resample); " - "canvas=Image.new(\"RGB\", img.size, (0, 0, 0)); " - "canvas.paste(logo, ((img.width - logo.width)//2, (img.height - logo.height)//2)); " - "canvas.save(\"/tmp/bg.jpg\")" - ) - - -def patched_weston_bg_exec_line() -> str: - python_cmd = patched_weston_bg_python().replace('"', '\\"') - return f"ExecStartPre=/bin/bash -c \"/usr/local/venv/bin/python -c '{python_cmd}'\"" - - -def patch_comma_sh_display_wait(original: bytes) -> bytes: - text = original.decode("utf-8") - if COMMA_SH_DISPLAY_WAIT_PATCH_MARKER in text: - return original - - old = """echo "waiting for magic" -for i in {1..200}; do - if systemctl is-active --quiet magic && [ -S /tmp/drmfd.sock ]; then - break - fi - sleep 0.1 -done - -if systemctl is-active --quiet magic && [ -S /tmp/drmfd.sock ]; then - echo "magic ready after ${SECONDS}s" -else - echo "timed out waiting for magic, ${SECONDS}s" -fi -""" - new = f"""# {COMMA_SH_DISPLAY_WAIT_PATCH_MARKER} -if systemctl cat magic.service >/dev/null 2>&1; then - echo "waiting for magic" - for i in {{1..200}}; do - if systemctl is-active --quiet magic && [ -S /tmp/drmfd.sock ]; then - break - fi - sleep 0.1 - done - - if systemctl is-active --quiet magic && [ -S /tmp/drmfd.sock ]; then - echo "magic ready after ${{SECONDS}}s" - else - echo "timed out waiting for magic, ${{SECONDS}}s" - fi -else - echo "magic unavailable; waiting for weston" - for i in {{1..200}}; do - if systemctl is-active --quiet weston-ready && [ -S /var/tmp/weston/wayland-0 ]; then - break - fi - sleep 0.1 - done - - if systemctl is-active --quiet weston-ready && [ -S /var/tmp/weston/wayland-0 ]; then - echo "weston ready after ${{SECONDS}}s" - else - echo "timed out waiting for weston, ${{SECONDS}}s" - fi -fi -""" - - if old not in text: - raise RuntimeError("Unable to find comma.sh display readiness wait") - return text.replace(old, new, 1).encode("utf-8") - - -def patch_weston_service(original: bytes) -> bytes: - text = original.decode("utf-8") - if WESTON_BG_PATCH_MARKER in text: - return original - - old = ( - "ExecStartPre=/bin/bash -c \"/usr/local/venv/bin/python -c 'from PIL import Image; " - "img=Image.open(\\\"/usr/comma/bg.jpg\\\"); " - "(img.rotate(180) if open(\\\"/sys/devices/platform/vendor/vendor:gpio-som-id/som_id\\\").read().strip() == \\\"1\\\" else img).save(\\\"/tmp/bg.jpg\\\")'\"" - ) - new = ( - f"# {WESTON_BG_PATCH_MARKER}: displayed boot logo was 90 degrees counter-clockwise.\n" - f"{patched_weston_bg_exec_line()}" - ) - - if old not in text: - raise RuntimeError("Unable to find weston.service background image generation line") - return text.replace(old, new, 1).encode("utf-8") - - -def patch_setup_branding_script(original: bytes, entry_name: str) -> bytes: - text = original.decode("utf-8") - if SETUP_BRANDING_PATCH_MARKER in text: - return text.encode("utf-8") - - text = text.replace( - 'OPENPILOT_URL = "https://openpilot.comma.ai"', - 'NETWORK_CHECK_URL = "https://openpilot.comma.ai"\n' - 'DEFAULT_INSTALLER_URL = "https://installer.comma.ai/firestar5683/StarPilot"\n' - f'# {SETUP_BRANDING_PATCH_MARKER}', - ) - text = text.replace("urllib.request.Request(OPENPILOT_URL, method=\"HEAD\")", - "urllib.request.Request(NETWORK_CHECK_URL, method=\"HEAD\")") - text = text.replace("urllib.request.urlopen(OPENPILOT_URL, timeout=2)", - "urllib.request.urlopen(NETWORK_CHECK_URL, timeout=2)") - text = text.replace("self.download(OPENPILOT_URL)", "self.download(DEFAULT_INSTALLER_URL)") - - if entry_name == MICI_SETUP_ENTRY_IN_SETUP_ZIPAPP: - text = text.replace('LargerSlider("slide to use\\nopenpilot"', 'LargerSlider("slide to use\\nstarpilot"') - text = text.replace('LargerSlider("slide to install\\nopenpilot"', 'LargerSlider("slide to install\\nstarpilot"') - text = text.replace('BigPillButton("install openpilot"', 'BigPillButton("install StarPilot"') - text = text.replace('set_text("install openpilot"', 'set_text("install StarPilot"') - elif entry_name == TICI_SETUP_ENTRY_IN_SETUP_ZIPAPP: - text = text.replace('ButtonRadio("openpilot"', 'ButtonRadio("StarPilot"') - - if SETUP_BRANDING_PATCH_MARKER not in text: - raise RuntimeError(f"Failed to patch setup branding for {entry_name}") - - return text.encode("utf-8") - - -def patch_setup_module(relative_path: str) -> bytes: - """ - Replace setup zipapp module with repo version so setup behavior stays in sync. - """ - repo_root = Path(__file__).resolve().parents[2] - src = repo_root / relative_path - if not src.is_file(): - raise RuntimeError(f"Unable to find repo setup source: {src}") - return src.read_bytes() - - -def patch_setup_script_with_ssh_restore(relative_path: str) -> bytes: - """ - Apply SSH-key restore logic directly into setup scripts for AGNOS images. - This keeps repo runtime behavior unchanged while making image reset flows - resilient to setup/install failures. - """ - text = patch_setup_module(relative_path).decode("utf-8") - if SETUP_SSH_RESTORE_PATCH_MARKER in text: - return text.encode("utf-8") - - restore_block = f""" -# {SETUP_SSH_RESTORE_PATCH_MARKER} -def _restore_ssh_after_reset(): - backup_dir = "/cache/reset_backup" - params_dir = "/data/params/d" - if not os.path.isdir(backup_dir): - return - - restored = False - try: - os.makedirs(params_dir, exist_ok=True) - for key in ("GithubSshKeys", "SshEnabled"): - src = f"{{backup_dir}}/{{key}}" - dst = f"{{params_dir}}/{{key}}" - if not os.path.isfile(src): - continue - shutil.copyfile(src, dst) - os.chmod(dst, 0o600) - restored = True - - if restored: - os.system("sudo chown -R comma:comma /data/params >/dev/null 2>&1 || true") - os.system("sudo /usr/comma/set_ssh.sh >/tmp/setup_ssh_restore.log 2>&1 || true") - finally: - os.system(f"sudo rm -rf {{backup_dir}} >/dev/null 2>&1 || true") -""" - - if "def main():" not in text: - raise RuntimeError(f"Unable to patch setup script without main(): {relative_path}") - - text = text.replace("\ndef main():", f"{restore_block}\n\ndef main():", 1) - text = text.replace(" try:\n gui_app.init_window(", " try:\n _restore_ssh_after_reset()\n gui_app.init_window(", 1) - return text.encode("utf-8") - - -def get_setup_replacements() -> dict[str, bytes]: - """ - Keep the reference setup bundle intact and patch only the networking backend. - - The reference AGNOS setup bundle already contains the correct small-screen - selector and matching mici UI modules. Replacing those modules with repo-head - versions caused bootstrap incompatibilities. The only setup-side changes we - still need are the jeepney fallback plus the StarPilot branding/url strings. - """ - return { - WIFI_MANAGER_ENTRY_IN_SETUP_ZIPAPP: patch_setup_wifi_manager(), - } - - -def patch_updater_module() -> bytes: - """ - Replace only the bundled updater selector with the repo version. - - The selector itself carries the small-screen fallback logic; the rest of the - reference updater zipapp stays unchanged. - """ - repo_root = Path(__file__).resolve().parents[2] - src = repo_root / "system/ui/updater.py" - if not src.is_file(): - raise RuntimeError(f"Unable to find repo updater source: {src}") - return src.read_bytes() - - -def patch_reset_script() -> bytes: - """ - Use repo reset.py so AGNOS reset stays in sync with upstream selector logic. - """ - repo_root = Path(__file__).resolve().parents[2] - src = repo_root / "system/ui/reset.py" - if not src.is_file(): - raise RuntimeError(f"Unable to find repo reset source: {src}") - data = src.read_text(encoding="utf-8") - if PATCH_MARKER not in data: - if data.startswith("#!"): - first_nl = data.find("\n") - if first_nl != -1: - data = data[:first_nl + 1] + f"# {PATCH_MARKER}\n" + data[first_nl + 1:] - else: - data = data + f"\n# {PATCH_MARKER}\n" - else: - data = f"# {PATCH_MARKER}\n" + data - return data.encode("utf-8") - def parse_args() -> argparse.Namespace: - p = argparse.ArgumentParser(description="Patch AGNOS system image with StarPilot reset and hardware support") - p.add_argument("--manifest", default="system/hardware/tici/agnos.json", help="Path to AGNOS manifest JSON") - p.add_argument("--work-dir", default=".cache/agnos_reset_patch", help="Working directory") - p.add_argument("--source-url", default=None, help="Override source raw system image URL") - p.add_argument("--source-image", default=None, help="Use existing local raw system image file instead of download") - p.add_argument("--reference-manifest", default=None, help="Optional AGNOS manifest used to source /usr/comma installer payloads") - p.add_argument("--reference-source-url", default=None, help="Override reference AGNOS system image URL for /usr/comma file sync") - p.add_argument("--reference-image", default=None, help="Use existing local reference system image file for /usr/comma file sync") - p.add_argument("--sync-comma-files", default=",".join(DEFAULT_SYNC_COMMA_FILES), - help="Comma-separated file list to sync from reference image (e.g. /usr/comma/installer,/usr/comma/setup)") - p.add_argument("--disable-comma-file-sync", action="store_true", - help="Disable syncing /usr/comma files from a reference image") - p.add_argument("--disable-usbgpu-firmware", action="store_true", - help="Do not install the AMD firmware required by the external GPU") - p.add_argument("--output-xz", default=None, help="Output .img.xz path") - p.add_argument("--new-url", default=None, help="Hosted URL for patched image; used for manifest output") - p.add_argument("--manifest-out", default=None, help="Write updated manifest JSON here") - p.add_argument("--in-place-manifest", action="store_true", help="Update manifest file in place") - p.add_argument("--force-download", action="store_true", help="Force redownload source image") - p.add_argument("--set-version", default=None, help="Override /VERSION inside patched image (e.g. 12.8.1)") - return p.parse_args() + parser = argparse.ArgumentParser(description="Build an upstream-based StarPilot AGNOS system image") + parser.add_argument("--manifest", default="system/hardware/tici/agnos.json", + help="Manifest to copy when writing an optional candidate manifest") + parser.add_argument("--source-url", default=UPSTREAM_SYSTEM_URL, help="Exact upstream system image URL") + parser.add_argument("--source-image", help="Use a local exact upstream raw, sparse, or .xz image") + parser.add_argument("--c3-deps-url", default=C3_DEPENDENCY_SOURCE_URL, + help="Exact prior StarPilot image containing the compatibility packages") + parser.add_argument("--c3-deps-image", help="Use a local exact StarPilot dependency source image") + parser.add_argument("--set-version", required=True, help="StarPilot revision, for example 19.6.5") + parser.add_argument("--work-dir", default=".cache/agnos_upstream_system") + parser.add_argument("--output-xz", help="Output .img.xz path") + parser.add_argument("--new-url", help="Hosted output URL for an optional candidate manifest") + parser.add_argument("--manifest-out", help="Candidate manifest output path; never overwrites the checked-in manifest") + parser.add_argument("--force-download", action="store_true") + return parser.parse_args() + + +def find_tool(name: str, extra_candidates: tuple[str, ...] = ()) -> str: + for candidate in (os.environ.get(name.upper()), name, *extra_candidates): + if candidate and (shutil.which(candidate) or Path(candidate).is_file()): + return candidate + raise RuntimeError(f"{name} not found") def find_debugfs() -> str: - candidates = [ - os.environ.get("DEBUGFS"), - "debugfs", - "/opt/homebrew/opt/e2fsprogs/sbin/debugfs", - ] - for c in candidates: - if c and shutil.which(c): - return c - if c and Path(c).is_file(): - return c - raise RuntimeError("debugfs not found. Install e2fsprogs and retry.") + return find_tool("debugfs", ("/opt/homebrew/opt/e2fsprogs/sbin/debugfs",)) -def load_manifest(path: Path) -> list[dict]: - return json.loads(path.read_text()) +def find_e2fsck() -> str: + return find_tool("e2fsck", ("/opt/homebrew/opt/e2fsprogs/sbin/e2fsck",)) -def get_system_entry(manifest: list[dict]) -> dict: - for e in manifest: - if e.get("name") == "system": - return e - raise RuntimeError("No system entry found in manifest") - - -def pick_source_url(system_entry: dict, override: str | None) -> str: - if override: - return override - url = system_entry.get("url") - if isinstance(url, str): - return url - alt = system_entry.get("alt") - if isinstance(alt, dict) and isinstance(alt.get("url"), str): - return alt["url"] - raise RuntimeError("No source URL found for system image") - - -def find_default_reference_manifest(primary_manifest_path: Path) -> Path | None: - # Expected tree layout for local development: - # /starpilot/system/hardware/tici/agnos.json - # /openpilot/system/hardware/tici/agnos.json - repo_root = primary_manifest_path - for _ in range(4): - if repo_root.parent == repo_root: - break - repo_root = repo_root.parent - - candidates = [ - repo_root.parent / "openpilot/openpilot/system/hardware/tici/agnos.json", - repo_root.parent / "openpilot/system/hardware/tici/agnos.json", - repo_root / "openpilot/system/hardware/tici/agnos.json", - ] - - for candidate in candidates: - if candidate.is_file() and candidate.resolve() != primary_manifest_path.resolve(): - return candidate.resolve() - return None - - -def parse_sync_file_list(raw: str) -> list[str]: - out: list[str] = [] - seen: set[str] = set() - for token in raw.replace(";", ",").split(","): - item = token.strip() - if not item: - continue - if not item.startswith("/"): - if "/" in item: - item = f"/{item.lstrip('/')}" - else: - item = f"/usr/comma/{item}" - if item not in seen: - seen.add(item) - out.append(item) - return out - - -def download(url: str, dst: Path) -> None: - dst.parent.mkdir(parents=True, exist_ok=True) - tmp = dst.with_suffix(dst.suffix + ".part") - print(f"Downloading {url} -> {dst}", flush=True) - with urllib.request.urlopen(url) as src, open(tmp, "wb") as out: - shutil.copyfileobj(src, out, length=1024 * 1024) - tmp.replace(dst) - - -def download_with_sha256(url: str, dst: Path, expected_sha256: str) -> None: - if not dst.exists(): - download(url, dst) - - actual_sha256 = sha256_file(dst) - if actual_sha256 != expected_sha256: - dst.unlink(missing_ok=True) - download(url, dst) - actual_sha256 = sha256_file(dst) - - if actual_sha256 != expected_sha256: - raise RuntimeError(f"Downloaded file hash mismatch for {dst}: got {actual_sha256}, expected {expected_sha256}") - - -def run_cmd(cmd: list[str], check: bool = True, capture: bool = True) -> subprocess.CompletedProcess[str]: - proc = subprocess.run(cmd, text=True, capture_output=capture) - if check and proc.returncode != 0: - raise RuntimeError(f"Command failed ({proc.returncode}): {' '.join(cmd)}\n{proc.stdout}\n{proc.stderr}") - return proc - - -def is_xz_file(path: Path) -> bool: - with open(path, "rb") as f: - header = f.read(len(XZ_MAGIC)) - return header == XZ_MAGIC - - -def decompress_xz(src: Path, dst: Path) -> None: - dst.parent.mkdir(parents=True, exist_ok=True) - tmp = dst.with_suffix(dst.suffix + ".part") - print(f"Decompressing XZ image {src} -> {dst}", flush=True) - with open(tmp, "wb") as out: - proc = subprocess.run(["xz", "-d", "-c", str(src)], stdout=out, stderr=subprocess.PIPE, text=True) - if proc.returncode != 0: - tmp.unlink(missing_ok=True) - raise RuntimeError(f"xz decompression failed:\n{proc.stderr}") - tmp.replace(dst) - - -def is_android_sparse(path: Path) -> bool: - with open(path, "rb") as f: - header = f.read(4) - if len(header) != 4: - return False - return int.from_bytes(header, "little") == ANDROID_SPARSE_MAGIC - - -def unsparse_image(src_sparse: Path, dst_raw: Path) -> None: - print(f"Unsparsing Android image {src_sparse} -> {dst_raw}", flush=True) - with open(src_sparse, "rb") as f_in, open(dst_raw, "wb") as f_out: - file_hdr = f_in.read(28) - if len(file_hdr) != 28: - raise RuntimeError("Invalid sparse image header length") - magic, major, minor, file_hdr_sz, chunk_hdr_sz, blk_sz, total_blks, total_chunks, _checksum = struct.unpack(" 28: - f_in.read(file_hdr_sz - 28) - if chunk_hdr_sz < 12: - raise RuntimeError(f"Invalid chunk header size: {chunk_hdr_sz}") - - for _ in range(total_chunks): - chunk_hdr = f_in.read(12) - if len(chunk_hdr) != 12: - raise RuntimeError("Unexpected EOF in chunk header") - chunk_type, _reserved, chunk_sz, total_sz = struct.unpack("<2H2I", chunk_hdr) - if chunk_hdr_sz > 12: - f_in.read(chunk_hdr_sz - 12) - - data_sz = total_sz - chunk_hdr_sz - out_chunk_bytes = chunk_sz * blk_sz - - if chunk_type == CHUNK_TYPE_RAW: - if data_sz != out_chunk_bytes: - raise RuntimeError(f"RAW chunk size mismatch: data={data_sz} out={out_chunk_bytes}") - remaining = data_sz - while remaining > 0: - chunk = f_in.read(min(8 * 1024 * 1024, remaining)) - if not chunk: - raise RuntimeError("Unexpected EOF in RAW chunk") - f_out.write(chunk) - remaining -= len(chunk) - elif chunk_type == CHUNK_TYPE_FILL: - if data_sz != 4: - raise RuntimeError(f"FILL chunk expected 4 bytes, got {data_sz}") - pattern = f_in.read(4) - if len(pattern) != 4: - raise RuntimeError("Unexpected EOF in FILL chunk") - # Write as sparse hole if fill is zero for speed. - if pattern == b"\x00\x00\x00\x00": - f_out.seek(out_chunk_bytes, os.SEEK_CUR) - else: - unit = pattern * (blk_sz // 4) - for _ in range(chunk_sz): - f_out.write(unit) - elif chunk_type == CHUNK_TYPE_DONT_CARE: - if data_sz > 0: - f_in.read(data_sz) - f_out.seek(out_chunk_bytes, os.SEEK_CUR) - elif chunk_type == CHUNK_TYPE_CRC32: - if data_sz != 4: - raise RuntimeError(f"CRC32 chunk expected 4 bytes, got {data_sz}") - f_in.read(4) - else: - raise RuntimeError(f"Unknown sparse chunk type: 0x{chunk_type:04x}") - - f_out.truncate(total_blks * blk_sz) - - -def materialize_ext4_image(source_img: Path, raw_img: Path, work_dir: Path, label: str, force: bool = False) -> None: - source_for_sparse = source_img - - if is_xz_file(source_img): - decompressed = work_dir / f"{label}.decompressed.img" - if force and decompressed.exists(): - decompressed.unlink() - if not decompressed.exists(): - decompress_xz(source_img, decompressed) - source_for_sparse = decompressed - - if force and raw_img.exists(): - raw_img.unlink() - - if raw_img.exists(): - return - - if is_android_sparse(source_for_sparse): - unsparse_image(source_for_sparse, raw_img) - else: - shutil.copy2(source_for_sparse, raw_img) - - -def run_debugfs(debugfs: str, image: Path, request: str, write: bool = False) -> str: - cmd = [debugfs] - if write: - cmd.append("-w") - cmd += ["-R", request, str(image)] - proc = run_cmd(cmd, check=True, capture=True) - return f"{proc.stdout}\n{proc.stderr}" - - -def split_shebang(data: bytes) -> tuple[bytes, bytes]: - if data.startswith(b"#!"): - idx = data.find(b"\n") - if idx != -1: - return data[:idx + 1], data[idx + 1:] - return b"", data - - -def patch_reset_zipapp(original: bytes) -> bytes: - shebang, zip_payload = split_shebang(original) - - src_io = BytesIO(zip_payload) - dst_io = BytesIO() - changed = False - - replacement_reset = patch_reset_script() - reset_replacements = { - RESET_ENTRY_IN_ZIPAPP: replacement_reset, - } - - with zipfile.ZipFile(src_io, "r") as src, zipfile.ZipFile(dst_io, "w", compression=zipfile.ZIP_DEFLATED) as dst: - if APPLICATION_ENTRY_IN_ZIPAPP not in src.namelist(): - raise RuntimeError(f"{APPLICATION_ENTRY_IN_ZIPAPP} not found in reset zipapp") - - seen_entries: set[str] = set() - for info in src.infolist(): - seen_entries.add(info.filename) - payload = src.read(info.filename) - if info.filename in reset_replacements: - replacement_payload = reset_replacements[info.filename] - if payload != replacement_payload: - payload = replacement_payload - changed = True - elif info.filename == APPLICATION_ENTRY_IN_ZIPAPP: - patched_payload = patch_application_script(payload) - if patched_payload != payload: - payload = patched_payload - changed = True - - new_info = zipfile.ZipInfo(info.filename, info.date_time) - new_info.compress_type = zipfile.ZIP_DEFLATED - new_info.external_attr = info.external_attr - new_info.create_system = info.create_system - dst.writestr(new_info, payload) - - # Some reference images may not include reset.py; add missing entry explicitly. - default_external_attr = 0o100644 << 16 - for entry, payload in reset_replacements.items(): - if entry in seen_entries: - continue - new_info = zipfile.ZipInfo(entry) - new_info.compress_type = zipfile.ZIP_DEFLATED - new_info.external_attr = default_external_attr - new_info.create_system = 3 - dst.writestr(new_info, payload) - changed = True - - if not changed: - return original - return shebang + dst_io.getvalue() - - -def patch_updater_zipapp(original: bytes) -> bytes: - shebang, zip_payload = split_shebang(original) - - replacement = patch_updater_module() - src_io = BytesIO(zip_payload) - dst_io = BytesIO() - changed = False - found_updater = False - - with zipfile.ZipFile(src_io, "r") as src, zipfile.ZipFile(dst_io, "w", compression=zipfile.ZIP_DEFLATED) as dst: - for info in src.infolist(): - payload = src.read(info.filename) - if info.filename == UPDATER_ENTRY_IN_ZIPAPP: - found_updater = True - if payload != replacement: - payload = replacement - changed = True - - new_info = zipfile.ZipInfo(info.filename, info.date_time) - new_info.compress_type = zipfile.ZIP_DEFLATED - new_info.external_attr = info.external_attr - new_info.create_system = info.create_system - dst.writestr(new_info, payload) - - if not found_updater: - new_info = zipfile.ZipInfo(UPDATER_ENTRY_IN_ZIPAPP) - new_info.compress_type = zipfile.ZIP_DEFLATED - new_info.external_attr = 0o100644 << 16 - new_info.create_system = 3 - dst.writestr(new_info, replacement) - changed = True - - if not changed: - return original - return shebang + dst_io.getvalue() - - -def patch_setup_zipapp(original: bytes) -> bytes: - shebang, zip_payload = split_shebang(original) - - src_io = BytesIO(zip_payload) - dst_io = BytesIO() - changed = False - - replacements = get_setup_replacements() - - with zipfile.ZipFile(src_io, "r") as src, zipfile.ZipFile(dst_io, "w", compression=zipfile.ZIP_DEFLATED) as dst: - seen_entries: set[str] = set() - for info in src.infolist(): - seen_entries.add(info.filename) - payload = src.read(info.filename) - if info.filename in replacements and payload != replacements[info.filename]: - payload = replacements[info.filename] - changed = True - elif info.filename in (MICI_SETUP_ENTRY_IN_SETUP_ZIPAPP, TICI_SETUP_ENTRY_IN_SETUP_ZIPAPP): - patched_payload = patch_setup_branding_script(payload, info.filename) - if patched_payload != payload: - payload = patched_payload - changed = True - - new_info = zipfile.ZipInfo(info.filename, info.date_time) - new_info.compress_type = zipfile.ZIP_DEFLATED - new_info.external_attr = info.external_attr - new_info.create_system = info.create_system - dst.writestr(new_info, payload) - - default_external_attr = 0o100644 << 16 - for entry, payload in replacements.items(): - if entry in seen_entries: - continue - new_info = zipfile.ZipInfo(entry) - new_info.compress_type = zipfile.ZIP_DEFLATED - new_info.external_attr = default_external_attr - new_info.create_system = 3 - dst.writestr(new_info, payload) - changed = True - - if not changed: - return original - return shebang + dst_io.getvalue() - - -def zipapp_has_markers(data: bytes) -> bool: - _shebang, zip_payload = split_shebang(data) - with zipfile.ZipFile(BytesIO(zip_payload), "r") as z: - reset_script = z.read(RESET_ENTRY_IN_ZIPAPP) - tici_reset_script = z.read(TICI_RESET_ENTRY_IN_ZIPAPP) - mici_reset_script = z.read(MICI_RESET_ENTRY_IN_ZIPAPP) - app_script = z.read(APPLICATION_ENTRY_IN_ZIPAPP) - return ( - PATCH_MARKER.encode() in reset_script - and b"_device_tree_device_type" in reset_script - and b"gui_app.big_ui()" not in reset_script - and b"mici_setup" not in mici_reset_script - and b"jeepney" not in mici_reset_script - and b"mici_setup" not in tici_reset_script - and b"jeepney" not in tici_reset_script - and APP_PATCH_MARKER.encode() in app_script - ) - - -def setup_zipapp_has_expected_content(data: bytes) -> bool: - _shebang, zip_payload = split_shebang(data) - replacements = get_setup_replacements() - with zipfile.ZipFile(BytesIO(zip_payload), "r") as z: - try: - wifi_manager = z.read(WIFI_MANAGER_ENTRY_IN_SETUP_ZIPAPP) - if SETUP_WIFI_PATCH_MARKER.encode() not in wifi_manager: - return False - for entry, payload in replacements.items(): - if z.read(entry) != payload: - return False - for entry in (MICI_SETUP_ENTRY_IN_SETUP_ZIPAPP, TICI_SETUP_ENTRY_IN_SETUP_ZIPAPP): - setup_script = z.read(entry) - if SETUP_BRANDING_PATCH_MARKER.encode() not in setup_script: - return False - if b"installer.comma.ai/firestar5683/StarPilot" not in setup_script: - return False - mici_setup = z.read(MICI_SETUP_ENTRY_IN_SETUP_ZIPAPP) - if b"install openpilot" in mici_setup or b"slide to install\\nopenpilot" in mici_setup: - return False - except KeyError: - return False - return True - - -def updater_zipapp_has_expected_content(data: bytes) -> bool: - _shebang, zip_payload = split_shebang(data) - with zipfile.ZipFile(BytesIO(zip_payload), "r") as z: - try: - updater_script = z.read(UPDATER_ENTRY_IN_ZIPAPP) - except KeyError: - return False - return updater_script == patch_updater_module() - - -def weston_service_has_expected_content(data: bytes) -> bool: - return ( - WESTON_BG_PATCH_MARKER.encode("utf-8") in data - and b"displayed boot logo was 90 degrees counter-clockwise" in data - and b"logo=img.crop(bbox) if bbox else img" in data - and b"logo=logo.rotate(-90, expand=True)" in data - and b"logo=logo.resize((max(1, logo.width//3), max(1, logo.height//3)), resample)" in data - and b"canvas.save(\\\"/tmp/bg.jpg\\\")" in data - ) - - -def comma_sh_has_expected_display_wait(data: bytes) -> bool: - return ( - COMMA_SH_DISPLAY_WAIT_PATCH_MARKER.encode("utf-8") in data - and b"systemctl cat magic.service" in data - and b"systemctl is-active --quiet weston-ready" in data - and b"[ -S /var/tmp/weston/wayland-0 ]" in data - ) - - -def parse_inode(debugfs_output: str) -> int: - m = re.search(r"Inode:\s+(\d+)", debugfs_output) - if not m: - raise RuntimeError(f"Unable to parse inode from debugfs stat output:\n{debugfs_output}") - return int(m.group(1)) - - -def format_debugfs_mode(mode_octal: str) -> str: - try: - mode = int(mode_octal, 8) - except ValueError as e: - raise RuntimeError(f"Invalid octal inode mode: {mode_octal}") from e - if not 0 <= mode <= 0xFFFF: - raise RuntimeError(f"Inode mode exceeds ext4 field width: {mode_octal}") - return f"0{mode:o}" - - -def verify_inode_metadata(debugfs: str, image: Path, image_path: str, expected_type: str, - mode_octal: str, uid: int, gid: int) -> None: - stat_out = run_debugfs(debugfs, image, f"stat {image_path}", write=False) - file_type, perms_octal, actual_uid, actual_gid = parse_debugfs_stat(stat_out) - expected_perms = int(mode_octal, 8) & 0o7777 - actual_perms = int(perms_octal, 8) - if (file_type, actual_perms, actual_uid, actual_gid) != (expected_type, expected_perms, uid, gid): - raise RuntimeError( - f"Metadata verification failed for {image_path}: " - f"got type={file_type} mode={actual_perms:04o} uid={actual_uid} gid={actual_gid}, " - f"expected type={expected_type} mode={expected_perms:04o} uid={uid} gid={gid}" - ) - - -def write_regular_file_to_image(debugfs: str, image: Path, image_path: str, local_file: Path, mode_octal: str, uid: int = 0, gid: int = 0) -> None: - try: - run_debugfs(debugfs, image, f"rm {image_path}", write=True) - except Exception as e: - err = str(e).lower() - if "file not found" not in err and "no such file" not in err: - raise - run_debugfs(debugfs, image, f"write {local_file} {image_path}", write=True) - stat_out = run_debugfs(debugfs, image, f"stat {image_path}", write=False) - inode = parse_inode(stat_out) - run_debugfs(debugfs, image, f"set_inode_field <{inode}> mode {format_debugfs_mode(mode_octal)}", write=True) - run_debugfs(debugfs, image, f"set_inode_field <{inode}> uid {uid}", write=True) - run_debugfs(debugfs, image, f"set_inode_field <{inode}> gid {gid}", write=True) - verify_inode_metadata(debugfs, image, image_path, "regular", mode_octal, uid, gid) - - -def ensure_directory_in_image(debugfs: str, image: Path, image_path: str, mode_octal: str = "040755", uid: int = 0, gid: int = 0) -> None: - try: - run_debugfs(debugfs, image, f"mkdir {image_path}", write=True) - except Exception as e: - err = str(e).lower() - if "already exists" not in err and "file exists" not in err: - raise - - stat_out = run_debugfs(debugfs, image, f"stat {image_path}", write=False) - inode = parse_inode(stat_out) - run_debugfs(debugfs, image, f"set_inode_field <{inode}> mode {format_debugfs_mode(mode_octal)}", write=True) - run_debugfs(debugfs, image, f"set_inode_field <{inode}> uid {uid}", write=True) - run_debugfs(debugfs, image, f"set_inode_field <{inode}> gid {gid}", write=True) - verify_inode_metadata(debugfs, image, image_path, "directory", mode_octal, uid, gid) - - -def extract_wheel_subset(wheel_path: Path, extract_dir: Path, roots: set[str]) -> dict[Path, str]: - if extract_dir.exists(): - shutil.rmtree(extract_dir) - extract_dir.mkdir(parents=True) - file_modes: dict[Path, str] = {} - - with zipfile.ZipFile(wheel_path, "r") as wheel: - for info in wheel.infolist(): - parts = Path(info.filename).parts - if not parts or parts[0] not in roots: - continue - if any(part == ".." for part in parts): - raise RuntimeError(f"Unsafe wheel entry path: {info.filename}") - - local_path = extract_dir.joinpath(*parts) - if info.is_dir(): - local_path.mkdir(parents=True, exist_ok=True) - continue - - local_path.parent.mkdir(parents=True, exist_ok=True) - with wheel.open(info, "r") as src, open(local_path, "wb") as dst: - shutil.copyfileobj(src, dst) - - perms = (info.external_attr >> 16) & 0o777 - if not perms: - perms = 0o644 - os.chmod(local_path, perms) - file_modes[local_path] = f"100{perms:03o}" - - if not (extract_dir / JEEPNY_PACKAGE_DIR / "__init__.py").is_file(): - raise RuntimeError("jeepney wheel extraction did not produce jeepney/__init__.py") - if not (extract_dir / JEEPNY_DIST_INFO_DIR / "METADATA").is_file(): - raise RuntimeError("jeepney wheel extraction did not produce dist-info/METADATA") - - return file_modes - - -def install_python_package_tree(debugfs: str, image: Path, source_dir: Path, image_root: str, file_modes: dict[Path, str]) -> None: - dirs = sorted((p for p in source_dir.rglob("*") if p.is_dir()), key=lambda p: len(p.relative_to(source_dir).parts)) - for local_dir in dirs: - rel = local_dir.relative_to(source_dir).as_posix() - ensure_directory_in_image(debugfs, image, f"{image_root}/{rel}", "040755", 0, 0) - - files = sorted((p for p in source_dir.rglob("*") if p.is_file()), key=lambda p: p.relative_to(source_dir).as_posix()) - for local_file in files: - rel = local_file.relative_to(source_dir).as_posix() - mode_octal = file_modes.get(local_file, "100644") - write_regular_file_to_image(debugfs, image, f"{image_root}/{rel}", local_file, mode_octal, 0, 0) - - -def install_jeepney_into_image(debugfs: str, image: Path, work_dir: Path) -> None: - wheel_dir = work_dir / "python_wheels" - wheel_path = wheel_dir / f"jeepney-{JEEPNY_VERSION}-py3-none-any.whl" - download_with_sha256(JEEPNY_WHEEL_URL, wheel_path, JEEPNY_WHEEL_SHA256) - - extract_dir = work_dir / "jeepney_wheel" - file_modes = extract_wheel_subset(wheel_path, extract_dir, {JEEPNY_PACKAGE_DIR, JEEPNY_DIST_INFO_DIR}) - - print(f"Installing jeepney {JEEPNY_VERSION} into AGNOS Python venv", flush=True) - install_python_package_tree(debugfs, image, extract_dir, PYTHON_SITE_PACKAGES_PATH_IN_IMAGE, file_modes) - - -def image_has_jeepney(debugfs: str, image: Path, work_dir: Path) -> bool: - verify_dir = work_dir / "jeepney_verify" - verify_dir.mkdir(parents=True, exist_ok=True) - init_file = verify_dir / "__init__.py" - wrappers_file = verify_dir / "wrappers.py" - metadata_file = verify_dir / "METADATA" - init_file.unlink(missing_ok=True) - wrappers_file.unlink(missing_ok=True) - metadata_file.unlink(missing_ok=True) - try: - for image_path in ( - f"{PYTHON_SITE_PACKAGES_PATH_IN_IMAGE}/{JEEPNY_PACKAGE_DIR}/__init__.py", - f"{PYTHON_SITE_PACKAGES_PATH_IN_IMAGE}/{JEEPNY_PACKAGE_DIR}/wrappers.py", - f"{PYTHON_SITE_PACKAGES_PATH_IN_IMAGE}/{JEEPNY_DIST_INFO_DIR}/METADATA", - ): - file_type, _mode, _uid, _gid = parse_debugfs_stat(run_debugfs(debugfs, image, f"stat {image_path}", write=False)) - if file_type != "regular": - raise RuntimeError(f"{image_path} has inode type {file_type}, expected regular") - - run_debugfs(debugfs, image, f"dump -p {PYTHON_SITE_PACKAGES_PATH_IN_IMAGE}/{JEEPNY_PACKAGE_DIR}/__init__.py {init_file}", write=False) - run_debugfs(debugfs, image, f"dump -p {PYTHON_SITE_PACKAGES_PATH_IN_IMAGE}/{JEEPNY_PACKAGE_DIR}/wrappers.py {wrappers_file}", write=False) - run_debugfs(debugfs, image, f"dump -p {PYTHON_SITE_PACKAGES_PATH_IN_IMAGE}/{JEEPNY_DIST_INFO_DIR}/METADATA {metadata_file}", write=False) - except Exception: - return False - - return ( - b"from .wrappers import *" in init_file.read_bytes() - and b"class DBusAddress" in wrappers_file.read_bytes() - and f"Version: {JEEPNY_VERSION}".encode("utf-8") in metadata_file.read_bytes() - ) - - -def parse_debugfs_stat(debugfs_output: str) -> tuple[str, str, int, int]: - type_match = re.search(r"Type:\s+([A-Za-z]+)", debugfs_output) - mode_match = re.search(r"Mode:\s+([0-7]+)", debugfs_output) - user_match = re.search(r"User:\s+(\d+)", debugfs_output) - group_match = re.search(r"Group:\s+(\d+)", debugfs_output) - if not type_match or not mode_match or not user_match or not group_match: - raise RuntimeError(f"Unable to parse debugfs stat output:\n{debugfs_output}") - return type_match.group(1).lower(), mode_match.group(1), int(user_match.group(1)), int(group_match.group(1)) - - -def inode_mode_from_type_and_perms(file_type: str, perms_octal: str) -> str: - prefix = INODE_MODE_TYPE_PREFIX.get(file_type) - if prefix is None: - raise RuntimeError(f"Unsupported inode type '{file_type}' for mode conversion") - perms = perms_octal.strip() - if not perms: - raise RuntimeError("Empty permissions value in inode stat") - return f"{prefix}{int(perms, 8):03o}" - - -def sync_files_from_reference_image(debugfs: str, reference_img: Path, patched_img: Path, sync_paths: list[str], work_dir: Path) -> list[str]: - sync_dir = work_dir / "reference_sync" - sync_dir.mkdir(parents=True, exist_ok=True) - synced: list[str] = [] - - for image_path in sync_paths: - source_tmp = sync_dir / f"source{image_path.replace('/', '_')}" - verify_tmp = sync_dir / f"verify{image_path.replace('/', '_')}" - - stat_out = run_debugfs(debugfs, reference_img, f"stat {image_path}", write=False) - file_type, perms_octal, uid, gid = parse_debugfs_stat(stat_out) - mode_octal = inode_mode_from_type_and_perms(file_type, perms_octal) - - run_debugfs(debugfs, reference_img, f"dump -p {image_path} {source_tmp}", write=False) - print(f"Syncing {image_path} from reference image (mode={mode_octal}, uid={uid}, gid={gid})", flush=True) - write_regular_file_to_image(debugfs, patched_img, image_path, source_tmp, mode_octal, uid, gid) - - run_debugfs(debugfs, patched_img, f"dump -p {image_path} {verify_tmp}", write=False) - if sha256_file(source_tmp) != sha256_file(verify_tmp): - raise RuntimeError(f"Verification failed after syncing {image_path}") - synced.append(image_path) - - return synced +def run_cmd(command: list[str], *, allowed_returncodes: frozenset[int] = frozenset({0})) -> subprocess.CompletedProcess[str]: + result = subprocess.run(command, check=False, capture_output=True, text=True) + if result.returncode not in allowed_returncodes: + raise RuntimeError(f"Command failed ({result.returncode}): {' '.join(command)}\n{result.stdout}\n{result.stderr}") + return result def sha256_file(path: Path) -> str: - h = hashlib.sha256() - with open(path, "rb") as f: - while True: - chunk = f.read(1024 * 1024) - if not chunk: - break - h.update(chunk) - return h.hexdigest() - - -def sha256_zstd_payload(path: Path) -> str: - try: - import zstandard - except ImportError as e: - raise RuntimeError("zstandard is required to verify the external-GPU firmware") from e - digest = hashlib.sha256() - with open(path, "rb") as compressed: - with zstandard.ZstdDecompressor().stream_reader(compressed) as source: - while chunk := source.read(1024 * 1024): - digest.update(chunk) + with path.open("rb") as stream: + while chunk := stream.read(8 * 1024 * 1024): + digest.update(chunk) return digest.hexdigest() -def install_amdgpu_firmware_from_reference(debugfs: str, reference_img: Path, patched_img: Path, work_dir: Path) -> None: - ensure_directory_in_image(debugfs, patched_img, AMDGPU_FIRMWARE_PATH_IN_IMAGE, "040755", 0, 0) - firmware_dir = work_dir / "amdgpu_firmware" - firmware_dir.mkdir(parents=True, exist_ok=True) - - for filename, expected_payload_hash in AMDGPU_FIRMWARE_SHA256.items(): - image_path = f"{AMDGPU_FIRMWARE_PATH_IN_IMAGE}/{filename}" - local_file = firmware_dir / filename - verify_file = firmware_dir / f"{filename}.verify" - local_file.unlink(missing_ok=True) - verify_file.unlink(missing_ok=True) - - stat_out = run_debugfs(debugfs, reference_img, f"stat {image_path}", write=False) - file_type, perms_octal, uid, gid = parse_debugfs_stat(stat_out) - if file_type != "regular": - raise RuntimeError(f"Reference firmware {image_path} is {file_type}, expected regular") - run_debugfs(debugfs, reference_img, f"dump -p {image_path} {local_file}", write=False) - if sha256_zstd_payload(local_file) != expected_payload_hash: - raise RuntimeError(f"Reference firmware payload hash mismatch for {filename}") - - mode_octal = inode_mode_from_type_and_perms(file_type, perms_octal) - write_regular_file_to_image(debugfs, patched_img, image_path, local_file, mode_octal, uid, gid) - run_debugfs(debugfs, patched_img, f"dump -p {image_path} {verify_file}", write=False) - if sha256_file(local_file) != sha256_file(verify_file): - raise RuntimeError(f"Compressed firmware verification failed for {filename}") - if sha256_zstd_payload(verify_file) != expected_payload_hash: - raise RuntimeError(f"Installed firmware payload hash mismatch for {filename}") - - print(f"Installed and verified {len(AMDGPU_FIRMWARE_SHA256)} AMD firmware files", flush=True) +def replace_exactly(text: str, old: str, new: str, expected_count: int = 1) -> str: + count = text.count(old) + if count != expected_count: + raise RuntimeError(f"Expected {expected_count} occurrences of {old!r}, found {count}") + return text.replace(old, new) -def compress_xz(src: Path, dst: Path) -> None: - dst.parent.mkdir(parents=True, exist_ok=True) - tmp = dst.with_suffix(dst.suffix + ".part") - print(f"Compressing {src} -> {dst}", flush=True) - with open(tmp, "wb") as out: - proc = subprocess.run(["xz", "-T0", "-6", "-c", str(src)], stdout=out, stderr=subprocess.PIPE, text=True) - if proc.returncode != 0: - tmp.unlink(missing_ok=True) - raise RuntimeError(f"xz failed:\n{proc.stderr}") - tmp.replace(dst) +def patch_setup_source(member: str, source: str) -> str: + bundled_installer_helper = r''' + +def patch_bundled_installer(path: str, owner: str, branch: str) -> None: + data = bytearray(open(path, "rb").read()) + + def patch_slot(old_value: bytes, new_value: bytes) -> None: + old_marker = old_value + b"?" + start = data.find(old_marker) + if start < 0 or data.find(old_marker, start + 1) >= 0: + raise RuntimeError(f"Expected exactly one installer slot for {old_value!r}") + end = data.find(b"\0", start) + if end < 0: + raise RuntimeError(f"Installer slot for {old_value!r} is not NUL terminated") + new_marker = new_value + b"?" + if len(new_marker) > end - start: + raise RuntimeError(f"Installer value {new_value!r} exceeds its slot") + data[start:end] = new_marker + b" " * (end - start - len(new_marker)) + + patch_slot(b"https://github.com/firestar5683/openpilot.git", f"https://github.com/{owner}/openpilot.git".encode("ascii")) + patch_slot(b"StarPilot", branch.encode("ascii")) + with open(path, "wb") as installer: + installer.write(data) -def update_manifest_system_entry(manifest: list[dict], new_url: str, new_hash_raw: str, size: int) -> list[dict]: - updated = json.loads(json.dumps(manifest)) - system_entry = get_system_entry(updated) - old_url = system_entry.get("url") - old_hash = system_entry.get("hash") - old_hash_raw = system_entry.get("hash_raw") - old_size = system_entry.get("size") +def install_bundled_installer(owner: str, branch: str, installer_url: str) -> None: + import tempfile - system_entry["url"] = new_url - system_entry["hash"] = new_hash_raw - system_entry["hash_raw"] = new_hash_raw - system_entry["size"] = size - system_entry["sparse"] = False - system_entry["full_check"] = False + fd, tmpfile = tempfile.mkstemp(prefix="installer_") + try: + with os.fdopen(fd, "wb") as destination, open("/usr/comma/installer", "rb") as source: + destination.write(source.read()) + patch_bundled_installer(tmpfile, owner, branch) + os.chmod(tmpfile, 0o755) + with open(INSTALLER_URL_PATH, "w") as installer_url_file: + installer_url_file.write(installer_url) + os.replace(tmpfile, INSTALLER_DESTINATION_PATH) + except Exception: + try: + os.close(fd) + except OSError: + pass + try: + os.unlink(tmpfile) + except FileNotFoundError: + pass + raise +''' + source = replace_exactly( + source, + 'OPENPILOT_URL = "https://openpilot.comma.ai"', + 'CONNECTIVITY_URL = "https://openpilot.comma.ai"\nOPENPILOT_URL = "file:///usr/comma/installer"' + bundled_installer_helper, + ) - if isinstance(old_url, str) and isinstance(old_hash, str) and isinstance(old_hash_raw, str) and isinstance(old_size, int): - system_entry["alt"] = { - "url": old_url, - "hash": old_hash, - "hash_raw": old_hash_raw, - "size": old_size, - } + source = replace_exactly( + source, + ''' # autocomplete incomplete URLs + if re.match("^([^/.]+)/([^/]+)$", url): + url = f"https://installer.comma.ai/{url}" - return updated + parsed = urlparse(url, scheme='https') + self.download_url = (urlparse(f"https://{url}") if not parsed.netloc else parsed).geturl()''', + ''' # owner/branch installs use the bundled COMMA/GBM installer. The + # installer.comma.ai binary targets Wayland and cannot run in this AGNOS. + self.installer_url = ("https://installer.comma.ai/firestar5683/StarPilot" if url == OPENPILOT_URL else url) + self.bundled_installer_target = (("firestar5683", "StarPilot") if url == OPENPILOT_URL else None) + match = re.fullmatch(r"(?:https://installer\\.comma\\.ai/)?([A-Za-z0-9_.-]+)/([A-Za-z0-9_.-]+)", url) + if match: + self.bundled_installer_target = match.groups() + self.installer_url = f"https://installer.comma.ai/{'/'.join(self.bundled_installer_target)}" + url = OPENPILOT_URL + parsed = urlparse(url, scheme='https') + self.download_url = (urlparse(f"https://{url}") if not parsed.netloc else parsed).geturl()''', + ) -def resolve_reference_source_image(args: argparse.Namespace, primary_manifest_path: Path, work_dir: Path) -> Path: - if args.reference_image: - ref_image = Path(args.reference_image).resolve() - if not ref_image.is_file(): - raise RuntimeError(f"Reference image not found: {ref_image}") - return ref_image + source = replace_exactly( + source, + ''' try: + import tempfile +''', + ''' try: + bundled_target = self.bundled_installer_target + if bundled_target is not None: + install_bundled_installer(*bundled_target, self.installer_url) + time.sleep(0.1) + gui_app.request_close() + return - reference_manifest_path: Path | None = None - if args.reference_manifest: - reference_manifest_path = Path(args.reference_manifest).resolve() - if not reference_manifest_path.is_file(): - raise RuntimeError(f"Reference manifest not found: {reference_manifest_path}") - else: - reference_manifest_path = find_default_reference_manifest(primary_manifest_path) + import tempfile +''', + ) - if args.reference_source_url: - reference_url = args.reference_source_url - elif reference_manifest_path is not None: - reference_manifest = load_manifest(reference_manifest_path) - reference_entry = get_system_entry(reference_manifest) - reference_url = pick_source_url(reference_entry, None) - print(f"Using reference AGNOS manifest: {reference_manifest_path}", flush=True) - else: - raise RuntimeError( - "No reference image source found. Set --reference-image, --reference-source-url, or --reference-manifest." + source = replace_exactly( + source, + ''' req = urllib.request.Request(self.download_url, headers=headers) + + with open(tmpfile, 'wb') as f, urllib.request.urlopen(req, timeout=30) as response: + total_size = int(response.headers.get('content-length', 0))''', + ''' response = (open("/usr/comma/installer", "rb") if self.download_url == OPENPILOT_URL else + urllib.request.urlopen(urllib.request.Request(self.download_url, headers=headers), timeout=30)) + + with open(tmpfile, 'wb') as f, response: + total_size = (os.path.getsize("/usr/comma/installer") if self.download_url == OPENPILOT_URL else + int(response.headers.get('content-length', 0)))''', + ) + + source = replace_exactly( + source, + ''' if not is_elf: +''', + ''' if is_elf and self.bundled_installer_target is not None: + patch_bundled_installer(tmpfile, *self.bundled_installer_target) + + if not is_elf: +''', + ) + + source = replace_exactly(source, "f.write(self.download_url)", "f.write(self.installer_url)") + + if member.endswith("mici_setup.py"): + source = replace_exactly(source, "urllib.request.Request(OPENPILOT_URL, method=\"HEAD\")", + "urllib.request.Request(CONNECTIVITY_URL, method=\"HEAD\")") + source = replace_exactly(source, 'LargerSlider("slide to install\\nopenpilot", use_openpilot_callback)', + 'LargerSlider("slide to install\\nStarPilot", use_openpilot_callback)') + source = replace_exactly(source, 'BigPillButton("install openpilot", green=True)', + 'BigPillButton("install StarPilot", green=True)') + source = replace_exactly(source, 'set_text("install openpilot" if not custom_software else "choose software")', + 'set_text("install StarPilot" if not custom_software else "choose software")') + source = replace_exactly(source, '"No custom software found at this URL: " + self.download_url.replace', + '"No custom software found at this URL: " + self.installer_url.replace') + source = replace_exactly( + source, + ''' except Exception: + self._download_failed_reason = "Invalid URL: " + self.download_url.replace("https://", "", 1)''', + ''' except Exception: + import traceback + traceback.print_exc() + self._download_failed_reason = "Invalid URL: " + self.installer_url.replace("https://", "", 1)''', ) + elif member.endswith("tici_setup.py"): + source = replace_exactly(source, "urllib.request.urlopen(OPENPILOT_URL, timeout=2.0)", + "urllib.request.urlopen(CONNECTIVITY_URL, timeout=2.0)") + source = replace_exactly(source, 'ButtonRadio("openpilot", self.checkmark', + 'ButtonRadio("StarPilot", self.checkmark') + source = replace_exactly(source, "self.download_failed(self.download_url,", "self.download_failed(self.installer_url,", expected_count=3) + source = replace_exactly( + source, + ''' except Exception: + error_msg = "Ensure the entered URL is valid, and the device's internet connection is good."''', + ''' except Exception: + import traceback + traceback.print_exc() + error_msg = "Ensure the entered URL is valid, and the device's internet connection is good."''', + ) + else: + raise RuntimeError(f"Unexpected setup source member: {member}") + return source - reference_download = work_dir / "reference_system.img" - if args.force_download and reference_download.exists(): - reference_download.unlink() - if not reference_download.exists(): - download(reference_url, reference_download) - return reference_download + +def is_setup_cache_member(member: str) -> bool: + return any( + member.startswith(str(Path(source_member).parent / "__pycache__" / Path(source_member).stem)) and member.endswith(".pyc") + for source_member in SETUP_SOURCE_MEMBERS + ) + + +def patch_setup_zipapp(source: Path, destination: Path) -> None: + raw = source.read_bytes() + zip_offset = raw.find(b"PK\x03\x04") + if zip_offset < 0: + raise RuntimeError("Setup payload is not an executable zip application") + prefix = raw[:zip_offset] + + destination.parent.mkdir(parents=True, exist_ok=True) + destination.write_bytes(prefix) + with zipfile.ZipFile(source, "r") as input_zip, zipfile.ZipFile(destination, "a") as output_zip: + names = set(input_zip.namelist()) + missing = set(SETUP_SOURCE_MEMBERS) - names + if missing: + raise RuntimeError(f"Setup payload is missing source members: {sorted(missing)}") + for info in input_zip.infolist(): + if is_setup_cache_member(info.filename): + continue + data = input_zip.read(info.filename) + if info.filename in SETUP_SOURCE_MEMBERS: + data = patch_setup_source(info.filename, data.decode("utf-8")).encode("utf-8") + output_zip.writestr(info, data) + + destination.chmod(source.stat().st_mode & 0o7777) + with zipfile.ZipFile(destination, "r") as patched_zip: + if patched_zip.testzip() is not None: + raise RuntimeError("Patched setup zip application failed CRC validation") + for member in SETUP_SOURCE_MEMBERS: + patched = patched_zip.read(member).decode("utf-8") + if 'OPENPILOT_URL = "file:///usr/comma/installer"' not in patched: + raise RuntimeError(f"Patched setup member does not use the bundled installer: {member}") + + +def patch_padded_binary_slot(data: bytearray, old_value: bytes, new_value: bytes) -> None: + old_marker = old_value + b"?" + start = data.find(old_marker) + if start < 0 or data.find(old_marker, start + 1) >= 0: + raise RuntimeError(f"Expected exactly one installer slot for {old_value!r}") + end = data.find(0, start) + if end < 0: + raise RuntimeError(f"Installer slot for {old_value!r} is not NUL terminated") + slot_length = end - start + new_marker = new_value + b"?" + if len(new_marker) > slot_length: + raise RuntimeError(f"Installer value {new_value!r} exceeds its {slot_length}-byte slot") + data[start:end] = new_marker + b" " * (slot_length - len(new_marker)) + + +def patch_installer_binary(source: Path, destination: Path) -> None: + data = bytearray(source.read_bytes()) + if not data.startswith(b"\x7fELF"): + raise RuntimeError("Bundled installer is not an ELF executable") + original_size = len(data) + patch_padded_binary_slot(data, b"https://github.com/commaai/openpilot.git", STAR_PILOT_GIT_URL.encode()) + patch_padded_binary_slot(data, b"release3", STAR_PILOT_BRANCH.encode()) + if len(data) != original_size: + raise RuntimeError("Installer patch changed the executable size") + destination.parent.mkdir(parents=True, exist_ok=True) + destination.write_bytes(data) + destination.chmod(source.stat().st_mode & 0o7777) + + +def validate_factory_install_payloads(setup: Path, installer: Path) -> None: + with zipfile.ZipFile(setup, "r") as setup_zip: + for member in SETUP_SOURCE_MEMBERS: + source = setup_zip.read(member).decode("utf-8") + if "StarPilot" not in source or 'OPENPILOT_URL = "file:///usr/comma/installer"' not in source: + raise RuntimeError(f"StarPilot setup customization is missing from {member}") + if "CONNECTIVITY_URL = \"https://openpilot.comma.ai\"" not in source: + raise RuntimeError(f"Connectivity check changed unexpectedly in {member}") + for expected in ( + "patch_bundled_installer(tmpfile, *self.bundled_installer_target)", + "install_bundled_installer(*bundled_target, self.installer_url)", + 'self.bundled_installer_target = (("firestar5683", "StarPilot") if url == OPENPILOT_URL else None)', + "self.bundled_installer_target = match.groups()", + "f.write(self.installer_url)", + ): + if expected not in source: + raise RuntimeError(f"Bundled custom-branch installer flow is missing from {member}: {expected}") + + installer_data = installer.read_bytes() + for expected in (STAR_PILOT_GIT_URL.encode() + b"?", STAR_PILOT_BRANCH.encode() + b"?"): + if installer_data.count(expected) != 1: + raise RuntimeError(f"Bundled installer is missing {expected!r}") + + +def validate_target_version(version: str) -> str: + clean = version.strip() + if not re.fullmatch(r"19\.6\.\d+", clean): + raise RuntimeError("Target version must be a 19.6.x StarPilot revision") + if int(clean.rsplit(".", 1)[1]) < 1: + raise RuntimeError("Target version must be newer than upstream 19.6") + return clean + + +def download(url: str, destination: Path) -> None: + destination.parent.mkdir(parents=True, exist_ok=True) + partial = destination.with_suffix(destination.suffix + ".part") + print(f"Downloading exact upstream AGNOS: {url}", flush=True) + with urllib.request.urlopen(url) as source, partial.open("wb") as output: + shutil.copyfileobj(source, output, length=8 * 1024 * 1024) + partial.replace(destination) + + +def is_xz_file(path: Path) -> bool: + with path.open("rb") as stream: + return stream.read(len(XZ_MAGIC)) == XZ_MAGIC + + +def decompress_xz(source: Path, destination: Path) -> None: + partial = destination.with_suffix(destination.suffix + ".part") + print(f"Decompressing {source}", flush=True) + with lzma.open(source, "rb") as compressed, partial.open("wb") as output: + shutil.copyfileobj(compressed, output, length=8 * 1024 * 1024) + partial.replace(destination) + + +def is_android_sparse(path: Path) -> bool: + with path.open("rb") as stream: + raw = stream.read(4) + return len(raw) == 4 and struct.unpack(" None: + print(f"Expanding Android sparse image {source}", flush=True) + with source.open("rb") as source_file, destination.open("wb") as output: + header = source_file.read(28) + if len(header) != 28: + raise RuntimeError("Sparse image header is truncated") + magic, major, _minor, file_header_size, chunk_header_size, block_size, total_blocks, total_chunks, _checksum = struct.unpack( + " None: + candidate = source + if is_xz_file(source): + decompressed = work_dir / "upstream_system.decompressed.img" + if not decompressed.exists(): + decompress_xz(source, decompressed) + candidate = decompressed + if destination.exists(): + return + if is_android_sparse(candidate): + unsparse_image(candidate, destination) + else: + shutil.copy2(candidate, destination) + + +def run_debugfs(debugfs: str, image: Path, request: str, *, write: bool = False) -> str: + command = [debugfs] + if write: + command.append("-w") + command += ["-R", request, str(image)] + result = run_cmd(command) + return f"{result.stdout}\n{result.stderr}" + + +def parse_inode(output: str) -> int: + match = re.search(r"Inode:\s+(\d+)", output) + if not match: + raise RuntimeError(f"Unable to parse inode:\n{output}") + return int(match.group(1)) + + +def write_version(debugfs: str, image: Path, local_file: Path) -> None: + expected_mutations = frozenset({ + VERSION_PATH_IN_IMAGE, + *STAR_PILOT_DEPENDENCY_PATHS, + *LEGACY_RUNTIME_LIBRARY_PATHS, + *FACTORY_INSTALL_PATHS, + }) + if ALLOWED_IMAGE_MUTATIONS != expected_mutations: + raise RuntimeError("AGNOS mutation allowlist contains an unexpected path") + run_debugfs(debugfs, image, f"rm {VERSION_PATH_IN_IMAGE}", write=True) + run_debugfs(debugfs, image, f"write {local_file} {VERSION_PATH_IN_IMAGE}", write=True) + inode = parse_inode(run_debugfs(debugfs, image, f"stat {VERSION_PATH_IN_IMAGE}")) + for field, value in (("mode", "0100644"), ("uid", "0"), ("gid", "0")): + run_debugfs(debugfs, image, f"set_inode_field <{inode}> {field} {value}", write=True) + + +def read_image_text(debugfs: str, image: Path, image_path: str) -> str: + output = run_debugfs(debugfs, image, f"cat {image_path}") + lines = [line.strip() for line in output.splitlines() if line.strip() and not line.startswith("debugfs ")] + return lines[0] if lines else "" + + +def image_path_exists(debugfs: str, image: Path, image_path: str) -> bool: + output = run_debugfs(debugfs, image, f"stat {image_path}") + return re.search(r"Inode:\s+\d+", output) is not None + + +def list_image_directory(debugfs: str, image: Path, image_path: str) -> set[str]: + output = run_debugfs(debugfs, image, f"ls -p {image_path}") + entries: set[str] = set() + for line in output.splitlines(): + if line.startswith("/"): + fields = line.split("/") + if len(fields) >= 6 and fields[5] not in ("", ".", ".."): + entries.add(fields[5]) + return entries + + +def validate_venv_layout(debugfs: str, image: Path, *, expected_count: int, + required_paths: dict[str, str]) -> dict[str, object]: + entries = list_image_directory(debugfs, image, SITE_PACKAGES_PATH_IN_IMAGE) + if len(entries) != expected_count: + raise RuntimeError(f"Managed venv has {len(entries)} entries; expected {expected_count}") + missing = [name for name, path in required_paths.items() if not image_path_exists(debugfs, image, path)] + if missing: + raise RuntimeError(f"Managed venv is missing required imports: {', '.join(missing)}") + return {"site_packages_count": len(entries), "required_imports": sorted(required_paths)} + + +def image_path_is_directory(debugfs: str, image: Path, image_path: str) -> bool: + output = run_debugfs(debugfs, image, f"stat {image_path}") + if not re.search(r"Inode:\s+\d+", output): + raise RuntimeError(f"Image path does not exist: {image_path}") + return "Type: directory" in output + + +def extract_starpilot_dependencies(debugfs: str, dependency_image: Path, destination: Path) -> None: + destination.mkdir(parents=True, exist_ok=True) + for image_path in STAR_PILOT_DEPENDENCY_PATHS: + if not image_path_exists(debugfs, dependency_image, image_path): + raise RuntimeError(f"StarPilot dependency source is missing {image_path}") + if image_path_is_directory(debugfs, dependency_image, image_path): + run_debugfs(debugfs, dependency_image, f"rdump {image_path} {destination}") + else: + run_debugfs(debugfs, dependency_image, f"dump -p {image_path} {destination / Path(image_path).name}") + extracted = {path.name for path in destination.iterdir()} + expected = set(STAR_PILOT_DEPENDENCY_NAMES) + if extracted != expected: + raise RuntimeError(f"Unexpected StarPilot dependency extraction: {sorted(extracted)}") + + +def extract_legacy_runtime_libraries(debugfs: str, dependency_image: Path, destination: Path) -> None: + destination.mkdir(parents=True, exist_ok=True) + for name, image_path in zip(LEGACY_RUNTIME_LIBRARY_NAMES, LEGACY_RUNTIME_LIBRARY_PATHS, strict=True): + stat = run_debugfs(debugfs, dependency_image, f"stat {image_path}") + if not re.search(r"Inode:\s+\d+", stat): + raise RuntimeError(f"Legacy runtime source is missing {image_path}") + local_path = destination / name + if "Type: symlink" in stat: + match = re.search(r'Fast link dest: "([^"]+)"', stat) + if match is None: + raise RuntimeError(f"Unable to resolve legacy runtime symlink {image_path}") + local_path.symlink_to(match.group(1)) + else: + run_debugfs(debugfs, dependency_image, f"dump -p {image_path} {local_path}") + + extracted = {path.name for path in destination.iterdir()} + if extracted != set(LEGACY_RUNTIME_LIBRARY_NAMES): + raise RuntimeError(f"Unexpected legacy runtime library extraction: {sorted(extracted)}") + for path in destination.iterdir(): + if path.is_symlink() and not path.resolve().is_file(): + raise RuntimeError(f"Legacy runtime symlink target is missing: {path}") + + +def ensure_image_directory(debugfs: str, image: Path, image_path: str) -> None: + if image_path_exists(debugfs, image, image_path): + return + parent = str(Path(image_path).parent) + if parent not in ("", ".", "/"): + ensure_image_directory(debugfs, image, parent) + run_debugfs(debugfs, image, f"mkdir {image_path}", write=True) + if not image_path_exists(debugfs, image, image_path): + raise RuntimeError(f"Failed to create image directory {image_path}") + inode = parse_inode(run_debugfs(debugfs, image, f"stat {image_path}")) + for field, value in (("mode", "040755"), ("uid", "0"), ("gid", "0")): + run_debugfs(debugfs, image, f"set_inode_field <{inode}> {field} {value}", write=True) + + +def add_tree_to_image(debugfs: str, image: Path, source: Path, destination: str) -> None: + if image_path_exists(debugfs, image, destination): + raise RuntimeError(f"Refusing to overwrite upstream image path {destination}") + commands: list[str] = [] + created_directories: set[str] = set() + + def create_directory(image_path: str) -> None: + if image_path in created_directories or image_path == SITE_PACKAGES_PATH_IN_IMAGE: + return + parent = str(Path(image_path).parent) + if parent not in ("", ".", "/"): + create_directory(parent) + commands.extend(( + f"mkdir {image_path}", + f"set_inode_field {image_path} mode 040755", + f"set_inode_field {image_path} uid 0", + f"set_inode_field {image_path} gid 0", + )) + created_directories.add(image_path) + + create_directory(destination) + for local_path in sorted(source.rglob("*")): + if "__pycache__" in local_path.parts: + continue + relative = local_path.relative_to(source) + image_path = str(Path(destination) / relative) + if local_path.is_dir(): + create_directory(image_path) + continue + create_directory(str(Path(image_path).parent)) + mode = "0100755" if local_path.stat().st_mode & 0o111 else "0100644" + commands.extend(( + f"write {local_path} {image_path}", + f"set_inode_field {image_path} mode {mode}", + f"set_inode_field {image_path} uid 0", + f"set_inode_field {image_path} gid 0", + )) + + with tempfile.NamedTemporaryFile("w", encoding="utf-8") as command_file: + command_file.write("\n".join(commands) + "\n") + command_file.flush() + result = run_cmd([debugfs, "-w", "-f", command_file.name, str(image)]) + output = f"{result.stdout}\n{result.stderr}" + if "File not found" in output or "Ext2 file already exists" in output: + raise RuntimeError(f"Failed to add StarPilot dependency tree {destination}:\n{output[-4000:]}") + if not image_path_exists(debugfs, image, destination): + raise RuntimeError(f"Failed to add StarPilot dependency tree {destination}") + + +def add_path_to_image(debugfs: str, image: Path, source: Path, destination: str) -> None: + if image_path_exists(debugfs, image, destination): + raise RuntimeError(f"Refusing to overwrite upstream image path {destination}") + if source.is_symlink(): + ensure_image_directory(debugfs, image, str(Path(destination).parent)) + target = os.readlink(source) + if "/" in target or target in ("", ".", ".."): + raise RuntimeError(f"Unsafe compatibility-library symlink target: {target!r}") + run_debugfs(debugfs, image, f"symlink {destination} {target}", write=True) + stat = run_debugfs(debugfs, image, f"stat {destination}") + if "Type: symlink" not in stat or f'Fast link dest: "{target}"' not in stat: + raise RuntimeError(f"Failed to add StarPilot dependency symlink {destination}") + return + if source.is_dir(): + add_tree_to_image(debugfs, image, source, destination) + return + if not source.is_file(): + raise RuntimeError(f"Extracted dependency path is missing: {source}") + ensure_image_directory(debugfs, image, str(Path(destination).parent)) + run_debugfs(debugfs, image, f"write {source} {destination}", write=True) + if not image_path_exists(debugfs, image, destination): + raise RuntimeError(f"Failed to add StarPilot dependency file {destination}") + inode = parse_inode(run_debugfs(debugfs, image, f"stat {destination}")) + mode = "0100755" if source.stat().st_mode & 0o111 else "0100644" + for field, value in (("mode", mode), ("uid", "0"), ("gid", "0")): + run_debugfs(debugfs, image, f"set_inode_field <{inode}> {field} {value}", write=True) + + +def replace_image_file(debugfs: str, image: Path, source: Path, destination: str) -> None: + if destination not in FACTORY_INSTALL_PATHS: + raise RuntimeError(f"Refusing to replace non-factory-install path {destination}") + if not image_path_exists(debugfs, image, destination): + raise RuntimeError(f"Factory-install path is missing from upstream image: {destination}") + if not source.is_file(): + raise RuntimeError(f"Replacement payload is missing: {source}") + run_debugfs(debugfs, image, f"rm {destination}", write=True) + run_debugfs(debugfs, image, f"write {source} {destination}", write=True) + inode = parse_inode(run_debugfs(debugfs, image, f"stat {destination}")) + for field, value in (("mode", "0100755"), ("uid", "0"), ("gid", "0")): + run_debugfs(debugfs, image, f"set_inode_field <{inode}> {field} {value}", write=True) + + +def fingerprint_image_paths(debugfs: str, image: Path, paths: tuple[str, ...], work_dir: Path, + label: str) -> dict[str, str]: + output_dir = work_dir / f"fingerprints_{label}" + output_dir.mkdir(parents=True, exist_ok=True) + fingerprints: dict[str, str] = {} + for image_path in paths: + local_path = output_dir / image_path.strip("/").replace("/", "_") + run_debugfs(debugfs, image, f"dump -p {image_path} {local_path}") + fingerprints[image_path] = sha256_file(local_path) + return fingerprints + + +def validate_protected_payloads(actual: dict[str, str]) -> None: + differences = { + path: {"actual": actual.get(path), "expected": expected} + for path, expected in PROTECTED_PAYLOAD_HASHES.items() + if actual.get(path) != expected + } + if differences: + raise RuntimeError(f"Image is not exact upstream AGNOS: {json.dumps(differences, sort_keys=True)}") + + +def validate_ext4(e2fsck: str, image: Path) -> None: + result = run_cmd([e2fsck, "-fn", str(image)], allowed_returncodes=frozenset({0, 1, 2})) + output = f"{result.stdout}\n{result.stderr}" + if "UNEXPECTED INCONSISTENCY" in output or "Filesystem still has errors" in output: + raise RuntimeError(f"ext4 validation failed:\n{output}") + + +def compress_xz(source: Path, destination: Path) -> None: + destination.parent.mkdir(parents=True, exist_ok=True) + partial = destination.with_suffix(destination.suffix + ".part") + print(f"Compressing {source} -> {destination}", flush=True) + with partial.open("wb") as output: + result = subprocess.run(["xz", "-T0", "-6", "-c", str(source)], stdout=output, stderr=subprocess.PIPE) + if result.returncode != 0: + partial.unlink(missing_ok=True) + raise RuntimeError(result.stderr.decode("utf-8", "replace")) + partial.replace(destination) + + +def get_system_entry(manifest: list[dict]) -> dict: + return next(entry for entry in manifest if entry.get("name") == "system") + + +def update_manifest_system_entry(manifest: list[dict], new_url: str, raw_hash: str, size: int) -> list[dict]: + updated = json.loads(json.dumps(manifest)) + entry = get_system_entry(updated) + entry.update({ + "url": new_url, + "hash": raw_hash, + "hash_raw": raw_hash, + "size": size, + "sparse": False, + "full_check": False, + "has_ab": True, + "ondevice_hash": raw_hash, + }) + entry.pop("alt", None) + entry.pop("casync_caibx", None) + entry.pop("casync_store", None) + return updated def main() -> int: args = parse_args() - debugfs = find_debugfs() - - manifest_path = Path(args.manifest).resolve() - manifest = load_manifest(manifest_path) - system_entry = get_system_entry(manifest) - + target_version = validate_target_version(args.set_version) + debugfs, e2fsck = find_debugfs(), find_e2fsck() work_dir = Path(args.work_dir).resolve() work_dir.mkdir(parents=True, exist_ok=True) if args.source_image: - downloaded_img = Path(args.source_image).resolve() - if not downloaded_img.is_file(): - raise RuntimeError(f"Source image not found: {downloaded_img}") + source = Path(args.source_image).resolve() else: - source_url = pick_source_url(system_entry, args.source_url) - downloaded_img = work_dir / "base_system.img" - if args.force_download and downloaded_img.exists(): - downloaded_img.unlink() - if not downloaded_img.exists(): - download(source_url, downloaded_img) + source = work_dir / "upstream_system.img.xz" + if args.force_download: + source.unlink(missing_ok=True) + if not source.exists(): + download(args.source_url, source) + if not source.is_file(): + raise RuntimeError(f"Source image not found: {source}") - raw_img = work_dir / "base_system.ext4.img" - materialize_ext4_image(downloaded_img, raw_img, work_dir, "base_system", force=args.force_download) - - patched_img = work_dir / "patched_system.ext4.img" - if patched_img.exists(): - patched_img.unlink() - print(f"Copying source image -> {patched_img}", flush=True) - shutil.copy2(raw_img, patched_img) - - sync_paths = [] if args.disable_comma_file_sync else parse_sync_file_list(args.sync_comma_files) - reference_raw = None - if sync_paths or not args.disable_usbgpu_firmware: - reference_source_img = resolve_reference_source_image(args, manifest_path, work_dir) - reference_raw = work_dir / "reference_system.ext4.img" - materialize_ext4_image(reference_source_img, reference_raw, work_dir, "reference_system", force=args.force_download) - - if sync_paths: - assert reference_raw is not None - print(f"Syncing /usr/comma payload files from reference image: {reference_raw}", flush=True) - synced_files = sync_files_from_reference_image(debugfs, reference_raw, patched_img, sync_paths, work_dir) - print(f"Synced {len(synced_files)} /usr/comma files from reference image", flush=True) - - if not args.disable_usbgpu_firmware: - assert reference_raw is not None - print(f"Installing external-GPU firmware from reference image: {reference_raw}", flush=True) - install_amdgpu_firmware_from_reference(debugfs, reference_raw, patched_img, work_dir) - - preserved_paths = { - RESET_PATH_IN_IMAGE: "comma_reset", - SETUP_PATH_IN_IMAGE: "comma_setup", - COMMA_SH_PATH_IN_IMAGE: "comma_sh", - MAGIC_PATH_IN_IMAGE: "comma_magic", - BG_PATH_IN_IMAGE: "comma_bg", + upstream_raw = work_dir / "upstream_system.ext4.img" + materialize_upstream_image(source, upstream_raw, work_dir) + if upstream_raw.stat().st_size != UPSTREAM_RAW_SIZE: + raise RuntimeError(f"Upstream raw size mismatch: {upstream_raw.stat().st_size}") + upstream_hash = sha256_file(upstream_raw) + if upstream_hash != UPSTREAM_RAW_SHA256: + raise RuntimeError(f"Upstream raw hash mismatch: {upstream_hash}") + validate_ext4(e2fsck, upstream_raw) + if read_image_text(debugfs, upstream_raw, VERSION_PATH_IN_IMAGE) != UPSTREAM_VERSION: + raise RuntimeError("Source image does not contain upstream /VERSION=19.6") + upstream_venv = validate_venv_layout( + debugfs, upstream_raw, + expected_count=UPSTREAM_SITE_PACKAGES_COUNT, + required_paths=UPSTREAM_REQUIRED_VENV_PATHS, + ) + additive_dependency_paths = (*STAR_PILOT_DEPENDENCY_PATHS, *LEGACY_RUNTIME_LIBRARY_PATHS) + unexpected_dependency_paths = [ + path for path in additive_dependency_paths if image_path_exists(debugfs, upstream_raw, path) + ] + if unexpected_dependency_paths: + raise RuntimeError(f"Dependency allowlist overlaps upstream AGNOS: {unexpected_dependency_paths}") + upstream_payloads = fingerprint_image_paths(debugfs, upstream_raw, tuple(UPSTREAM_PAYLOAD_HASHES), work_dir, "upstream") + upstream_differences = { + path: {"actual": upstream_payloads.get(path), "expected": expected} + for path, expected in UPSTREAM_PAYLOAD_HASHES.items() + if upstream_payloads.get(path) != expected } - expected_hashes: dict[str, str] = {} - for image_path, label in preserved_paths.items(): - preserved_file = work_dir / f"{label}.preserved" - print(f"Recording existing {image_path} for preservation", flush=True) - run_debugfs(debugfs, patched_img, f"dump -p {image_path} {preserved_file}", write=False) - expected_hashes[image_path] = sha256_file(preserved_file) + if upstream_differences: + raise RuntimeError(f"Source image is not exact upstream AGNOS: {json.dumps(upstream_differences, sort_keys=True)}") - original_updater = work_dir / "comma_updater.orig" - patched_updater = work_dir / "comma_updater.patched" - verify_updater = work_dir / "comma_updater.verify" - original_weston = work_dir / "weston_service.orig" - patched_weston = work_dir / "weston_service.patched" - verify_weston = work_dir / "weston_service.verify" - patched_comma_sh = work_dir / "comma_sh.patched" + upstream_fingerprint_dir = work_dir / "fingerprints_upstream" + customized_payload_dir = work_dir / "starpilot_factory_install" + if customized_payload_dir.exists(): + shutil.rmtree(customized_payload_dir) + customized_payload_dir.mkdir(parents=True) + customized_setup = customized_payload_dir / "setup" + customized_installer = customized_payload_dir / "installer" + patch_setup_zipapp(upstream_fingerprint_dir / "usr_comma_setup", customized_setup) + patch_installer_binary(upstream_fingerprint_dir / "usr_comma_installer", customized_installer) + validate_factory_install_payloads(customized_setup, customized_installer) - original_comma_sh = work_dir / "comma_sh.preserved" - comma_sh_patched_data = patch_comma_sh_display_wait(original_comma_sh.read_bytes()) - patched_comma_sh.write_bytes(comma_sh_patched_data) + c3_work_dir = work_dir / "c3_dependency_source" + c3_work_dir.mkdir(parents=True, exist_ok=True) + if args.c3_deps_image: + c3_source = Path(args.c3_deps_image).resolve() + else: + c3_source = c3_work_dir / "system.img.xz" + if args.force_download: + c3_source.unlink(missing_ok=True) + if not c3_source.exists(): + download(args.c3_deps_url, c3_source) + if not c3_source.is_file(): + raise RuntimeError(f"C3 dependency source image not found: {c3_source}") + c3_raw = c3_work_dir / "system.ext4.img" + materialize_upstream_image(c3_source, c3_raw, c3_work_dir) + if c3_raw.stat().st_size != C3_DEPENDENCY_SOURCE_RAW_SIZE: + raise RuntimeError(f"C3 dependency source size mismatch: {c3_raw.stat().st_size}") + c3_source_hash = sha256_file(c3_raw) + if c3_source_hash != C3_DEPENDENCY_SOURCE_RAW_SHA256: + raise RuntimeError(f"C3 dependency source hash mismatch: {c3_source_hash}") + dependency_packages_dir = work_dir / "starpilot_dependency_packages" + if dependency_packages_dir.exists(): + shutil.rmtree(dependency_packages_dir) + extract_starpilot_dependencies(debugfs, c3_raw, dependency_packages_dir) + legacy_runtime_dir = work_dir / "starpilot_legacy_runtime_libraries" + if legacy_runtime_dir.exists(): + shutil.rmtree(legacy_runtime_dir) + extract_legacy_runtime_libraries(debugfs, c3_raw, legacy_runtime_dir) - print("Writing patched /usr/comma/comma.sh display readiness wait", flush=True) - write_regular_file_to_image(debugfs, patched_img, COMMA_SH_PATH_IN_IMAGE, patched_comma_sh, "100775", 0, 0) - expected_hashes[COMMA_SH_PATH_IN_IMAGE] = sha256_file(patched_comma_sh) + candidate_raw = work_dir / f"starpilot_system_{target_version}.ext4.img" + candidate_raw.unlink(missing_ok=True) + shutil.copy2(upstream_raw, candidate_raw) + version_file = work_dir / "VERSION.starpilot" + version_file.write_text(target_version + "\n", encoding="utf-8") + write_version(debugfs, candidate_raw, version_file) + for image_path in STAR_PILOT_DEPENDENCY_PATHS: + add_path_to_image(debugfs, candidate_raw, dependency_packages_dir / Path(image_path).name, image_path) + for image_path in LEGACY_RUNTIME_LIBRARY_PATHS: + add_path_to_image(debugfs, candidate_raw, legacy_runtime_dir / Path(image_path).name, image_path) + replace_image_file(debugfs, candidate_raw, customized_setup, SETUP_PATH_IN_IMAGE) + replace_image_file(debugfs, candidate_raw, customized_installer, INSTALLER_PATH_IN_IMAGE) - print("Extracting weston.service from image", flush=True) - run_debugfs(debugfs, patched_img, f"dump -p {WESTON_SERVICE_PATH_IN_IMAGE} {original_weston}", write=False) + if read_image_text(debugfs, candidate_raw, VERSION_PATH_IN_IMAGE) != target_version: + raise RuntimeError("Failed to write the StarPilot AGNOS version marker") + candidate_venv = validate_venv_layout( + debugfs, candidate_raw, + expected_count=CANDIDATE_SITE_PACKAGES_COUNT, + required_paths=REQUIRED_VENV_PATHS, + ) + missing_legacy_runtime = [ + path for path in LEGACY_RUNTIME_LIBRARY_PATHS if not image_path_exists(debugfs, candidate_raw, path) + ] + if missing_legacy_runtime: + raise RuntimeError(f"Candidate image is missing legacy runtime libraries: {missing_legacy_runtime}") + candidate_payloads = fingerprint_image_paths(debugfs, candidate_raw, tuple(PROTECTED_PAYLOAD_HASHES), work_dir, "candidate") + validate_protected_payloads(candidate_payloads) + upstream_protected_payloads = {path: upstream_payloads[path] for path in PROTECTED_PAYLOAD_HASHES} + if candidate_payloads != upstream_protected_payloads: + raise RuntimeError("Protected upstream payloads changed") + candidate_factory_payloads = fingerprint_image_paths( + debugfs, candidate_raw, tuple(UPSTREAM_FACTORY_INSTALL_HASHES), work_dir, "candidate_factory_install", + ) + expected_factory_payloads = { + SETUP_PATH_IN_IMAGE: sha256_file(customized_setup), + INSTALLER_PATH_IN_IMAGE: sha256_file(customized_installer), + } + if candidate_factory_payloads != expected_factory_payloads: + raise RuntimeError("Factory-install payloads do not match the validated StarPilot replacements") + candidate_factory_dir = work_dir / "fingerprints_candidate_factory_install" + validate_factory_install_payloads( + candidate_factory_dir / "usr_comma_setup", + candidate_factory_dir / "usr_comma_installer", + ) + validate_ext4(e2fsck, candidate_raw) - weston_original_data = original_weston.read_bytes() - weston_patched_data = patch_weston_service(weston_original_data) - if weston_patched_data == weston_original_data: - print("weston.service already contains the expected boot-logo patch; continuing", flush=True) - patched_weston.write_bytes(weston_patched_data) + raw_hash = sha256_file(candidate_raw) + output_xz = Path(args.output_xz).resolve() if args.output_xz else work_dir / f"system-{raw_hash}.img.xz" + compress_xz(candidate_raw, output_xz) + metadata = { + "base_version": UPSTREAM_VERSION, + "base_raw_sha256": UPSTREAM_RAW_SHA256, + "target_version": target_version, + "allowed_image_mutations": sorted(ALLOWED_IMAGE_MUTATIONS), + "raw_sha256": raw_hash, + "raw_size": candidate_raw.stat().st_size, + "xz_sha256": sha256_file(output_xz), + "xz_size": output_xz.stat().st_size, + "upstream_venv_validation": upstream_venv, + "candidate_venv_validation": candidate_venv, + "c3_dependency_source_raw_sha256": c3_source_hash, + "starpilot_dependency_paths": list(STAR_PILOT_DEPENDENCY_PATHS), + "c3_dependency_paths": list(C3_DEPENDENCY_PATHS), + "legacy_runtime_library_paths": list(LEGACY_RUNTIME_LIBRARY_PATHS), + "protected_payloads": candidate_payloads, + "factory_install_payloads": candidate_factory_payloads, + "factory_reset_stack": ( + "upstream reset/network/updater/Magic unchanged; setup uses the bundled stock COMMA/GBM installer " + "for both the default StarPilot install and custom GitHub owner/branch installs" + ), + "device_validation_required": True, + } + metadata_path = Path(str(output_xz) + ".metadata.json") + metadata_path.write_text(json.dumps(metadata, indent=2) + "\n", encoding="utf-8") - print("Writing patched weston.service back into image", flush=True) - write_regular_file_to_image(debugfs, patched_img, WESTON_SERVICE_PATH_IN_IMAGE, patched_weston, "100644", 0, 0) - - run_debugfs(debugfs, patched_img, f"dump -p {WESTON_SERVICE_PATH_IN_IMAGE} {verify_weston}", write=False) - verify_weston_data = verify_weston.read_bytes() - if not weston_service_has_expected_content(verify_weston_data): - raise RuntimeError("weston.service verification failed after writing weston.service file into image") - - print("Extracting /usr/comma/updater from image", flush=True) - run_debugfs(debugfs, patched_img, f"dump -p {UPDATER_PATH_IN_IMAGE} {original_updater}", write=False) - - updater_original_data = original_updater.read_bytes() - updater_patched_data = patch_updater_zipapp(updater_original_data) - if updater_patched_data == updater_original_data: - print("Updater zipapp already contains the expected selector patch; continuing", flush=True) - patched_updater.write_bytes(updater_patched_data) - - print("Writing patched /usr/comma/updater back into image", flush=True) - write_regular_file_to_image(debugfs, patched_img, UPDATER_PATH_IN_IMAGE, patched_updater, "100775", 0, 0) - - run_debugfs(debugfs, patched_img, f"dump -p {UPDATER_PATH_IN_IMAGE} {verify_updater}", write=False) - verify_updater_data = verify_updater.read_bytes() - if not updater_zipapp_has_expected_content(verify_updater_data): - raise RuntimeError("Updater zipapp verification failed after writing updater file into image") - - install_jeepney_into_image(debugfs, patched_img, work_dir) - if not image_has_jeepney(debugfs, patched_img, work_dir): - raise RuntimeError("jeepney verification failed after installing package into image") - - for image_path, label in preserved_paths.items(): - verify_file = work_dir / f"{label}.verify" - run_debugfs(debugfs, patched_img, f"dump -p {image_path} {verify_file}", write=False) - if image_path == COMMA_SH_PATH_IN_IMAGE and not comma_sh_has_expected_display_wait(verify_file.read_bytes()): - raise RuntimeError("comma.sh display readiness verification failed") - if sha256_file(verify_file) != expected_hashes[image_path]: - raise RuntimeError(f"{image_path} does not match the expected generated payload") - - if args.set_version: - version_file = work_dir / "VERSION.patched" - version_file.write_text(args.set_version.strip() + "\n", encoding="utf-8") - print(f"Writing {VERSION_PATH_IN_IMAGE}={args.set_version.strip()}", flush=True) - write_regular_file_to_image(debugfs, patched_img, VERSION_PATH_IN_IMAGE, version_file, "100644", 0, 0) - version_raw = run_debugfs(debugfs, patched_img, f"cat {VERSION_PATH_IN_IMAGE}", write=False) - version_lines = [ln.strip() for ln in version_raw.splitlines() if ln.strip() and not ln.startswith("debugfs ")] - version_verify = version_lines[0] if version_lines else "" - if version_verify != args.set_version.strip(): - raise RuntimeError(f"/VERSION mismatch after patch: got '{version_verify}', expected '{args.set_version.strip()}'") - - raw_hash = sha256_file(patched_img) - raw_size = patched_img.stat().st_size - - default_name = f"system-{raw_hash}.img.xz" - output_xz = Path(args.output_xz).resolve() if args.output_xz else (work_dir / default_name) - compress_xz(patched_img, output_xz) - - print("") - print("Patched AGNOS system artifact ready:") - print(f" raw image: {patched_img}") - print(f" xz image: {output_xz}") - print(f" raw sha256/hash_raw: {raw_hash}") - print(f" size: {raw_size}") - print("") + print("Validated upstream-based StarPilot AGNOS artifact:") + print(f" target version: {target_version}") + print(f" raw image: {candidate_raw}") + print(f" xz image: {output_xz}") + print(f" raw sha256: {raw_hash}") + print(f" xz sha256: {metadata['xz_sha256']}") + print(f" metadata: {metadata_path}") + print(" only mutations: /VERSION, additive StarPilot runtime/C3 compatibility, and factory setup/installer branding") if args.new_url: - new_manifest = update_manifest_system_entry(manifest, args.new_url, raw_hash, raw_size) - out_path: Path - if args.in_place_manifest: - out_path = manifest_path - elif args.manifest_out: - out_path = Path(args.manifest_out).resolve() - else: - out_path = work_dir / "agnos.patched.json" - out_path.write_text(json.dumps(new_manifest, indent=2) + "\n") - print(f"Updated manifest written: {out_path}") - else: - print("No --new-url provided. Manifest not updated.") - print("Set system entry values to:") - print(json.dumps({ - "url": "", - "hash": raw_hash, - "hash_raw": raw_hash, - "size": raw_size, - "sparse": False, - "full_check": False, - "has_ab": True, - }, indent=2)) - + manifest_path = Path(args.manifest).resolve() + output_manifest = Path(args.manifest_out).resolve() if args.manifest_out else work_dir / "agnos.candidate.json" + if output_manifest == manifest_path: + raise RuntimeError("Refusing to overwrite the checked-in manifest") + manifest = json.loads(manifest_path.read_text(encoding="utf-8")) + output_manifest.write_text( + json.dumps(update_manifest_system_entry(manifest, args.new_url, raw_hash, candidate_raw.stat().st_size), indent=2) + "\n", + encoding="utf-8", + ) + print(f" candidate manifest: {output_manifest}") + elif args.manifest_out: + raise RuntimeError("--manifest-out requires --new-url") return 0 diff --git a/tools/agnos/test_patch_system_reset_image.py b/tools/agnos/test_patch_system_reset_image.py index c2742978a..edec1d23c 100644 --- a/tools/agnos/test_patch_system_reset_image.py +++ b/tools/agnos/test_patch_system_reset_image.py @@ -1,106 +1,238 @@ +import importlib.util +import json +import os +import zipfile from pathlib import Path -import runpy import pytest -from tools.agnos.patch_system_reset_image import ( - AMDGPU_FIRMWARE_SHA256, - COMMA_SH_DISPLAY_WAIT_PATCH_MARKER, - comma_sh_has_expected_display_wait, - find_default_reference_manifest, - format_debugfs_mode, - patch_comma_sh_display_wait, - patch_setup_branding_script, - sha256_zstd_payload, -) + +def _load_patch_module(): + path = Path(__file__).resolve().parent / "patch_system_reset_image.py" + spec = importlib.util.spec_from_file_location("patch_system_reset_image_under_test", path) + assert spec is not None and spec.loader is not None + module = importlib.util.module_from_spec(spec) + spec.loader.exec_module(module) + return module -ORIGINAL_DISPLAY_WAIT = b'''#!/usr/bin/env bash -echo "waiting for magic" -for i in {1..200}; do - if systemctl is-active --quiet magic && [ -S /tmp/drmfd.sock ]; then - break - fi - sleep 0.1 -done +patch_image = _load_patch_module() -if systemctl is-active --quiet magic && [ -S /tmp/drmfd.sock ]; then - echo "magic ready after ${SECONDS}s" -else - echo "timed out waiting for magic, ${SECONDS}s" -fi -exec /data/continue.sh +def test_only_version_runtime_packages_and_factory_install_payloads_are_mutable(): + assert patch_image.ALLOWED_IMAGE_MUTATIONS == { + patch_image.VERSION_PATH_IN_IMAGE, + *patch_image.STAR_PILOT_DEPENDENCY_PATHS, + *patch_image.LEGACY_RUNTIME_LIBRARY_PATHS, + patch_image.SETUP_PATH_IN_IMAGE, + patch_image.INSTALLER_PATH_IN_IMAGE, + } + assert set(patch_image.C3_DEPENDENCY_PATHS) < set(patch_image.STAR_PILOT_DEPENDENCY_PATHS) + assert len(patch_image.LEGACY_RUNTIME_LIBRARY_PATHS) == 10 + assert len(patch_image.LEGACY_RUNTIME_LIBRARY_PATHS) == len(set(patch_image.LEGACY_RUNTIME_LIBRARY_PATHS)) + + +def test_upstream_profile_covers_factory_reset_and_runtime_paths(): + assert patch_image.UPSTREAM_VERSION == "19.6" + assert patch_image.UPSTREAM_RAW_SHA256 == "5b6ce7965904a157fd3a134ccfcb854f9ca5c1cc2a26b7cb80a4fa4e1cc4aaa3" + assert set(patch_image.UPSTREAM_REQUIRED_VENV_PATHS) == {"capnp", "numpy", "Crypto", "tqdm", "raylib"} + assert set(patch_image.REQUIRED_VENV_PATHS) == { + "crcmod", "serial", "kaitaistruct", "cv2", "mapbox_earcut", "jsonrpc", "xattr", "onnx", + "aiohttp", "pyaudio", "capnp", "numpy", "Crypto", "tqdm", "raylib", + } + assert patch_image.CANDIDATE_SITE_PACKAGES_COUNT == ( + patch_image.UPSTREAM_SITE_PACKAGES_COUNT + len(patch_image.STAR_PILOT_DEPENDENCY_PATHS) + ) + assert len(patch_image.STAR_PILOT_DEPENDENCY_PATHS) == len(set(patch_image.STAR_PILOT_DEPENDENCY_PATHS)) + assert set(patch_image.PROTECTED_PAYLOAD_HASHES) >= { + "/etc/NetworkManager/NetworkManager.conf", + "/lib/systemd/system/NetworkManager.service", + "/usr/comma/updater", + "/usr/comma/reset", + "/usr/comma/comma.sh", + "/usr/comma/magic.py", + } + assert set(patch_image.UPSTREAM_FACTORY_INSTALL_HASHES) == {"/usr/comma/installer", "/usr/comma/setup"} + assert set(patch_image.PROTECTED_PAYLOAD_HASHES).isdisjoint(patch_image.UPSTREAM_FACTORY_INSTALL_HASHES) + + +@pytest.mark.parametrize("version", ["19.6.1", "19.6.5", "19.6.99"]) +def test_target_version_accepts_starpilot_revision(version): + assert patch_image.validate_target_version(version) == version + + +@pytest.mark.parametrize("version", ["19.6", "19.6.0", "19.7.1", "20.6.1", "latest"]) +def test_target_version_rejects_non_revision(version): + with pytest.raises(RuntimeError): + patch_image.validate_target_version(version) + + +def test_write_version_fails_closed_if_allowlist_changes(tmp_path, monkeypatch): + monkeypatch.setattr(patch_image, "ALLOWED_IMAGE_MUTATIONS", frozenset({"/VERSION", "/usr/comma/setup"})) + with pytest.raises(RuntimeError, match="allowlist"): + patch_image.write_version("debugfs", tmp_path / "system.img", tmp_path / "VERSION") + + +def test_add_path_to_image_preserves_runtime_library_symlink(tmp_path, monkeypatch): + target = tmp_path / "libavformat.so.58.29.100" + target.write_bytes(b"ELF") + source = tmp_path / "libavformat.so.58" + source.symlink_to(target.name) + calls = [] + + monkeypatch.setattr(patch_image, "image_path_exists", lambda *_args: False) + monkeypatch.setattr(patch_image, "ensure_image_directory", lambda *_args: None) + + def fake_run_debugfs(_debugfs, _image, request, *, write=False): + calls.append((request, write)) + if request.startswith("stat "): + return 'Inode: 1 Type: symlink\nFast link dest: "libavformat.so.58.29.100"' + return "" + + monkeypatch.setattr(patch_image, "run_debugfs", fake_run_debugfs) + patch_image.add_path_to_image("debugfs", tmp_path / "system.img", source, "/usr/local/lib/libavformat.so.58") + + assert ("symlink /usr/local/lib/libavformat.so.58 libavformat.so.58.29.100", True) in calls + assert not any(request.startswith("write ") for request, _write in calls) + + +def _setup_source(member: str) -> str: + connectivity = ( + 'request = urllib.request.Request(OPENPILOT_URL, method="HEAD")' + if member.endswith("mici_setup.py") + else "urllib.request.urlopen(OPENPILOT_URL, timeout=2.0)" + ) + labels = ( + 'LargerSlider("slide to install\\nopenpilot", use_openpilot_callback)\n' + 'BigPillButton("install openpilot", green=True)\n' + 'set_text("install openpilot" if not custom_software else "choose software")' + if member.endswith("mici_setup.py") + else 'ButtonRadio("openpilot", self.checkmark)' + ) + if member.endswith("mici_setup.py"): + not_elf = 'self._download_failed_reason = "No custom software found at this URL: " + self.download_url.replace("https://", "", 1)' + http_error = 'self._download_failed_reason = "http"' + generic_error = 'self._download_failed_reason = "Invalid URL: " + self.download_url.replace("https://", "", 1)' + else: + not_elf = 'self.download_failed(self.download_url, "No custom software found at this URL.")' + http_error = 'self.download_failed(self.download_url, "http")' + generic_error = ( + 'error_msg = "Ensure the entered URL is valid, and the device\'s internet connection is good."\n' + ' self.download_failed(self.download_url, error_msg)' + ) + return f'''OPENPILOT_URL = "https://openpilot.comma.ai" +{connectivity} +{labels} + def download(self, url: str): + # autocomplete incomplete URLs + if re.match("^([^/.]+)/([^/]+)$", url): + url = f"https://installer.comma.ai/{{url}}" + + parsed = urlparse(url, scheme='https') + self.download_url = (urlparse(f"https://{{url}}") if not parsed.netloc else parsed).geturl() + + try: + import tempfile + + headers = {{"User-Agent": "test"}} + req = urllib.request.Request(self.download_url, headers=headers) + + with open(tmpfile, 'wb') as f, urllib.request.urlopen(req, timeout=30) as response: + total_size = int(response.headers.get('content-length', 0)) + is_elf = True + if not is_elf: + {not_elf} + with open(INSTALLER_URL_PATH, "w") as f: + f.write(self.download_url) + except urllib.error.HTTPError as e: + {http_error} + except Exception: + {generic_error} ''' -def test_patch_comma_sh_display_wait_uses_available_display_service(): - patched = patch_comma_sh_display_wait(ORIGINAL_DISPLAY_WAIT) +def test_patch_setup_zipapp_preserves_prefix_and_custom_flow(tmp_path): + source = tmp_path / "setup" + source.write_bytes(b"#!/usr/bin/env python3\n") + with zipfile.ZipFile(source, "a") as setup_zip: + for member in patch_image.SETUP_SOURCE_MEMBERS: + setup_zip.writestr(member, _setup_source(member)) + cache_member = str(Path(member).parent / "__pycache__" / f"{Path(member).stem}.cpython-312.pyc") + setup_zip.writestr(cache_member, b"stale") + setup_zip.writestr("unchanged.txt", b"upstream") + os.chmod(source, 0o755) - assert COMMA_SH_DISPLAY_WAIT_PATCH_MARKER.encode() in patched - assert b"systemctl cat magic.service" in patched - assert b"systemctl is-active --quiet magic" in patched - assert b"systemctl is-active --quiet weston-ready" in patched - assert b"[ -S /var/tmp/weston/wayland-0 ]" in patched - assert comma_sh_has_expected_display_wait(patched) - assert patch_comma_sh_display_wait(patched) == patched + destination = tmp_path / "setup.patched" + patch_image.patch_setup_zipapp(source, destination) + + assert destination.read_bytes().startswith(b"#!/usr/bin/env python3\n") + assert destination.stat().st_mode & 0o777 == 0o755 + with zipfile.ZipFile(destination) as setup_zip: + assert setup_zip.read("unchanged.txt") == b"upstream" + assert not any(patch_image.is_setup_cache_member(name) for name in setup_zip.namelist()) + mici = setup_zip.read(patch_image.SETUP_SOURCE_MEMBERS[0]).decode() + tici = setup_zip.read(patch_image.SETUP_SOURCE_MEMBERS[1]).decode() + assert 'OPENPILOT_URL = "file:///usr/comma/installer"' in mici + assert 'LargerSlider("slide to install\\nStarPilot"' in mici + assert "urllib.request.Request(CONNECTIVITY_URL" in mici + assert 'ButtonRadio("StarPilot"' in tici + assert "urllib.request.urlopen(CONNECTIVITY_URL" in tici + for setup_source in (mici, tici): + assert "patch_bundled_installer(tmpfile, *self.bundled_installer_target)" in setup_source + assert "install_bundled_installer(*bundled_target, self.installer_url)" in setup_source + assert 'self.bundled_installer_target = (("firestar5683", "StarPilot") if url == OPENPILOT_URL else None)' in setup_source + assert "self.installer_url = (" in setup_source + assert "url = OPENPILOT_URL" in setup_source + assert "f.write(self.installer_url)" in setup_source + assert 'open("/usr/comma/installer", "rb")' in setup_source + assert "self.download_url == OPENPILOT_URL" in setup_source + assert 'url = f"https://installer.comma.ai/{url}"' not in setup_source -def test_patch_comma_sh_display_wait_rejects_unknown_layout(): - with pytest.raises(RuntimeError, match="display readiness wait"): - patch_comma_sh_display_wait(b"#!/usr/bin/env bash\nexec /data/continue.sh\n") +def test_patch_installer_binary_keeps_elf_layout_and_targets_starpilot(tmp_path): + source = tmp_path / "installer" + source.write_bytes( + b"\x7fELF" + + b"https://github.com/commaai/openpilot.git?" + b" " * 64 + b"\0" + + b"release3?" + b" " * 64 + b"\0tail" + ) + os.chmod(source, 0o755) + destination = tmp_path / "installer.patched" + + patch_image.patch_installer_binary(source, destination) + + assert destination.stat().st_size == source.stat().st_size + assert destination.stat().st_mode & 0o777 == 0o755 + data = destination.read_bytes() + assert data.count(b"https://github.com/firestar5683/openpilot.git?") == 1 + assert data.count(b"StarPilot?") == 1 + assert b"commaai/openpilot" not in data -@pytest.mark.parametrize("slider_text", ["slide to use", "slide to install"]) -def test_patch_mici_setup_branding_handles_old_and_new_labels(slider_text): - original = f'''OPENPILOT_URL = "https://openpilot.comma.ai" -self._openpilot_slider = LargerSlider("{slider_text}\\nopenpilot", callback) -self._continue_button = BigPillButton("install openpilot", green=True) -self._continue_button.set_text("install openpilot" if not custom_software else "choose software") -'''.encode() - - patched = patch_setup_branding_script(original, "openpilot/system/ui/mici_setup.py") - - assert b"installer.comma.ai/firestar5683/StarPilot" in patched - assert f"{slider_text}\\nstarpilot".encode() in patched - assert b"install StarPilot" in patched - assert b"install openpilot" not in patched +def test_update_manifest_changes_only_system_entry(): + original = [ + {"name": "boot", "url": "custom-boot", "hash": "boot-hash"}, + {"name": "system", "url": "old", "hash": "old", "alt": {"url": "old-alt"}}, + ] + updated = patch_image.update_manifest_system_entry(original, "hosted", "new-hash", 123) + assert updated[0] == original[0] + assert updated[1] == { + "name": "system", + "url": "hosted", + "hash": "new-hash", + "hash_raw": "new-hash", + "size": 123, + "sparse": False, + "full_check": False, + "has_ab": True, + "ondevice_hash": "new-hash", + } + assert json.dumps(original) -@pytest.mark.parametrize(("mode", "expected"), [ - ("100775", "0100775"), - ("100644", "0100644"), - ("040755", "040755"), - ("120777", "0120777"), -]) -def test_format_debugfs_mode(mode, expected): - assert format_debugfs_mode(mode) == expected - - -def test_external_gpu_firmware_matches_tinygrad_requirements(): - firmware_metadata = Path(__file__).resolve().parents[2] / "tinygrad/runtime/autogen/am/fw.py" - hashes = runpy.run_path(firmware_metadata)["hashes"] - expected = {filename.removesuffix(".zst"): digest for filename, digest in AMDGPU_FIRMWARE_SHA256.items()} - assert all(hashes[filename] == digest for filename, digest in expected.items()) - - -def test_zstd_payload_hash(tmp_path): - import hashlib - import zstandard - - payload = b"external GPU firmware payload" - compressed = tmp_path / "firmware.bin.zst" - compressed.write_bytes(zstandard.ZstdCompressor().compress(payload)) - - assert sha256_zstd_payload(compressed) == hashlib.sha256(payload).hexdigest() - - -def test_default_reference_manifest_uses_sibling_openpilot(tmp_path): - primary = tmp_path / "starpilot/system/hardware/tici/agnos.json" - reference = tmp_path / "openpilot/openpilot/system/hardware/tici/agnos.json" - primary.parent.mkdir(parents=True) - reference.parent.mkdir(parents=True) - primary.write_text("[]") - reference.write_text("[]") - - assert find_default_reference_manifest(primary) == reference.resolve() +def test_protected_payload_validation_reports_any_drift(): + patch_image.validate_protected_payloads(dict(patch_image.PROTECTED_PAYLOAD_HASHES)) + changed = dict(patch_image.PROTECTED_PAYLOAD_HASHES) + changed["/usr/comma/reset"] = "0" * 64 + with pytest.raises(RuntimeError, match="/usr/comma/reset"): + patch_image.validate_protected_payloads(changed) diff --git a/tools/agnos/validate_agnos_runtime.sh b/tools/agnos/validate_agnos_runtime.sh new file mode 100755 index 000000000..992db70c2 --- /dev/null +++ b/tools/agnos/validate_agnos_runtime.sh @@ -0,0 +1,90 @@ +#!/usr/bin/env bash +set -euo pipefail + +EXPECTED_VERSION="${1:-}" +REPO_ROOT="${2:-/data/openpilot}" +PYTHON_BIN="/usr/local/venv/bin/python3" +SITE_PACKAGES="/usr/local/venv/lib/python3.12/site-packages" + +fail() { + echo "[FAIL] $*" >&2 + exit 1 +} + +[[ -x "$PYTHON_BIN" ]] || fail "managed Python is missing: $PYTHON_BIN" +[[ -d "$SITE_PACKAGES" ]] || fail "site-packages is missing: $SITE_PACKAGES" + +actual_version="$(tr -d '\r\n'