diff --git a/SConstruct b/SConstruct index 3dd8c9337..7147ecdc1 100644 --- a/SConstruct +++ b/SConstruct @@ -1,5 +1,4 @@ import os -import importlib import shutil import subprocess import sys @@ -129,34 +128,6 @@ 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' @@ -299,8 +270,6 @@ env = Environment( ] + cflags + ccflags, CPPPATH=cpppath + [ - capnproto_include_dirs, - ffmpeg_include_dirs, "#", "#third_party/acados/include", "#third_party/acados/include/blasfeo/include", @@ -319,13 +288,11 @@ env = Environment( RANLIB=ranlib, LINKFLAGS=ldflags, - RPATH=rpath + ffmpeg_lib_dirs, + RPATH=rpath, CFLAGS=["-std=gnu11"] + cflags, CXXFLAGS=["-std=c++1z"] + cxxflags, LIBPATH=libpath + [ - capnproto_lib_dirs, - ffmpeg_lib_dirs, "#msgq_repo", "#third_party", "#selfdrive/pandad", @@ -404,15 +371,7 @@ SConscript(['opendbc_repo/SConscript'], exports={'env': env_swaglog}) SConscript(['cereal/SConscript']) Import('socketmaster', 'msgq') -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'] +messaging = [socketmaster, msgq, 'capnp', 'kj',] Export('messaging') diff --git a/launch_env.sh b/launch_env.sh index 56b378f72..5693bbf91 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.10" + export AGNOS_VERSION="19.6.2" fi if [ -z "$AGNOS_ACCEPTED_VERSIONS" ]; then diff --git a/pyproject.toml b/pyproject.toml index ec6813d98..8dce92875 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -29,11 +29,6 @@ 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 0867df47c..6177692d9 100644 --- a/selfdrive/ui/SConscript +++ b/selfdrive/ui/SConscript @@ -1,29 +1,15 @@ -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": - 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"] + raylib_libs += ["GLESv2", "wayland-client", "wayland-egl", "EGL"] else: - raylib_env['LIBPATH'] += [f'#third_party/raylib/{arch}/'] - raylib_libs = common + ["raylib", "GL"] + raylib_libs += ["GL"] release = "release3" installers = [ @@ -37,15 +23,11 @@ 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, inter_bold, inter_light], LIBS=raylib_libs) - assert installer[0].get_size() < 2500*1e3, installer[0].get_size() + 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() diff --git a/selfdrive/ui/installer/installer.cc b/selfdrive/ui/installer/installer.cc index b79e18329..3e03ac0ff 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 "raylib.h" +#include "third_party/raylib/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 a93bddf7c..3e2c65740 100644 --- a/system/hardware/tici/agnos.json +++ b/system/hardware/tici/agnos.json @@ -56,7 +56,7 @@ }, { "name": "boot", - "url": "https://files.firestar.link/x/ugiq4cqx08q7/boot9.img.xz", + "url": "https://www.dropbox.com/scl/fi/9l9io42qfx2shr9er5jqx/boot9.img.xz?rlkey=lbtxz862kxbvn3jn98ejitj8p&st=vr35tgz7&dl=1", "hash": "ab2eba0f96b2f48efa376330c3eb509158361adf3ad9c20f269ec92457aa841f", "hash_raw": "ab2eba0f96b2f48efa376330c3eb509158361adf3ad9c20f269ec92457aa841f", "size": 48343040, @@ -67,13 +67,13 @@ }, { "name": "system", - "url": "https://files.firestar.link/x/07530yj1jd6a/system18.img.xz", - "hash": "01c84930849f9be2bdbad5e9a8dda3a6fd2be95e81d1b3556574b18346934f49", - "hash_raw": "01c84930849f9be2bdbad5e9a8dda3a6fd2be95e81d1b3556574b18346934f49", + "url": "https://www.dropbox.com/scl/fi/pewhzpqzi3aewuiaffc6m/system10.img.xz?rlkey=olzrzulhs93zzghnjrskmdwxt&st=exnfk2oz&dl=1", + "hash": "ab395d4c963a908ab86709f1a6580a62dd24cb34ee71cf5fd5ec29d7d48d0e10", + "hash_raw": "ab395d4c963a908ab86709f1a6580a62dd24cb34ee71cf5fd5ec29d7d48d0e10", "size": 4718592000, "sparse": false, "full_check": false, "has_ab": true, - "ondevice_hash": "01c84930849f9be2bdbad5e9a8dda3a6fd2be95e81d1b3556574b18346934f49" + "ondevice_hash": "ab395d4c963a908ab86709f1a6580a62dd24cb34ee71cf5fd5ec29d7d48d0e10" } ] diff --git a/system/loggerd/tests/test_uploader.py b/system/loggerd/tests/test_uploader.py index aa3b13dc3..2e13d357a 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 clear_locks, main, UPLOAD_ATTR_NAME, UPLOAD_ATTR_VALUE +from openpilot.system.loggerd.uploader import 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,10 +36,6 @@ 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 06ed867a9..827dc8128 100755 --- a/system/loggerd/uploader.py +++ b/system/loggerd/uploader.py @@ -63,9 +63,6 @@ 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 034b2f4c4..431635fa1 100755 --- a/system/manager/manager.py +++ b/system/manager/manager.py @@ -1094,6 +1094,11 @@ 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 39e9d8291..434c917b0 100644 --- a/system/manager/process.py +++ b/system/manager/process.py @@ -634,7 +634,19 @@ class PythonProcess(ManagerProcess): self.launcher = launcher def prepare(self) -> None: - pass + 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) 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 888f1739a..fbe19c3dc 100644 --- a/system/webrtc/device/video.py +++ b/system/webrtc/device/video.py @@ -1,19 +1,17 @@ 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.realtime import DT_MDL from openpilot.common.params import Params +from openpilot.common.realtime import DT_MDL -# 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, @@ -21,15 +19,6 @@ 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", @@ -63,9 +52,6 @@ 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: @@ -82,7 +68,9 @@ class LiveStreamVideoStreamTrack(TiciVideoStreamTrack): async def recv(self): while True: - # while video is disabled, pause here without returning + if self.readyState != "live": + raise MediaStreamError + if not self.video_enabled: await asyncio.sleep(0.005) continue @@ -91,11 +79,18 @@ 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, block=False) + self.params.put("LivestreamRequestKeyframe", 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 EncodedVideoFrame(self._build_frame_data(msg), self._pts) + return packet + + def codec_preference(self) -> str | None: + return "H264" diff --git a/system/webrtc/tests/test_stream_session.py b/system/webrtc/tests/test_stream_session.py index 7077e9877..f8316d203 100644 --- a/system/webrtc/tests/test_stream_session.py +++ b/system/webrtc/tests/test_stream_session.py @@ -1,17 +1,20 @@ 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 pytest - -pytest.importorskip("libdatachannel", reason="the upstream WebRTC backend requires Python 3.12") - +import pyaudio 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: @@ -30,37 +33,40 @@ class TestStreamSession: expected_dict = {"type": "customReservedRawData0", "logMonoTime": 123, "valid": True, "data": "test"} expected_json = json.dumps(expected_dict).encode() - channel = mocker.Mock() - channel.is_open.return_value = True - proxy = CerealOutgoingMessageProxy(["customReservedRawData0"]) - - def mocked_update(_): - proxy.sm.update_msgs(0, [test_msg]) + channel = mocker.Mock(spec=RTCDataChannel) + mocked_submaster = messaging.SubMaster(["customReservedRawData0"]) + def mocked_update(t): + mocked_submaster.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"}, - {"type": "can", "data": [{"address": 0, "dat": "", "src": 0}]}, - {"type": "testJoystick", "data": {"axes": [0, 0], "buttons": [False]}}, + {"type": "customReservedRawData0", "data": "test"}, # primitive + {"type": "can", "data": [{"address": 0, "dat": "", "src": 0}]}, # list + {"type": "testJoystick", "data": {"axes": [0, 0], "buttons": [False]}}, # dict ] 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() - msg_type, message = mocked_pubmaster.send.call_args.args - assert msg_type == msg["type"] - assert isinstance(message, capnp._DynamicStructBuilder) - assert hasattr(message, msg_type) + mt, md = mocked_pubmaster.send.call_args.args + assert mt == msg["type"] + assert isinstance(md, capnp._DynamicStructBuilder) + assert hasattr(md, msg["type"]) + mocked_pubmaster.reset_mock() def test_livestream_track(self, mocker): @@ -72,11 +78,29 @@ 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 - assert bytes(packet) == b"" + 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 diff --git a/system/webrtc/tests/test_webrtcd.py b/system/webrtc/tests/test_webrtcd.py index 9fb6a42e5..23a7f6ddc 100644 --- a/system/webrtc/tests/test_webrtcd.py +++ b/system/webrtc/tests/test_webrtcd.py @@ -1,42 +1,65 @@ -import json - 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 -pytest.importorskip("libdatachannel", reason="the upstream WebRTC backend requires Python 3.12") +from openpilot.system.webrtc.webrtcd import get_stream -from openpilot.system.webrtc.webrtcd import ServerState, handle_get_schema, handle_post_notify, on_shutdown +import aiortc +from teleoprtc import WebRTCOfferBuilder +from parameterized import parameterized_class +@parameterized_class(("in_services", "out_services"), [ + (["testJoystick"], ["carState"]), + ([], ["carState"]), + (["testJoystick"], []), + ([], []), +]) @pytest.mark.asyncio -async def test_get_schema(): - status, body, content_type = await handle_get_schema(ServerState(), "carState") +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") - assert status == 200 - assert content_type.startswith("application/json") - assert "carState" in json.loads(body) + 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) + 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() -@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") + stream = builder.stream() + await self.assertCompletesWithTimeout(stream.start()) + await self.assertCompletesWithTimeout(stream.wait_for_connection()) -@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 + 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) - status, body, content_type = await handle_post_notify(state, {"type": "ping"}) + 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()) - 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"})) + await self.assertCompletesWithTimeout(stream.stop()) - await on_shutdown(state) - - session.stop.assert_awaited_once() - assert state.streams == {} + # 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()) diff --git a/system/webrtc/webrtcd.py b/system/webrtc/webrtcd.py index b3bacc81e..837eda817 100644 --- a/system/webrtc/webrtcd.py +++ b/system/webrtc/webrtcd.py @@ -1,29 +1,33 @@ #!/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 -import signal -import threading -from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer -from urllib.parse import urlparse, parse_qs -from typing import Any +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 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 @@ -37,8 +41,20 @@ 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") @@ -67,10 +83,10 @@ class CerealOutgoingMessageProxy(AsyncTaskRunner): super().__init__() self.services = list(services) self.sm = messaging.SubMaster(self.services) - self.channels = [] + self.channels: list[RTCDataChannel] = [] self._enabled = enabled - def add_channel(self, channel): + def add_channel(self, channel: 'RTCDataChannel'): self.channels.append(channel) def enable(self, enable: bool): @@ -99,17 +115,20 @@ 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) @@ -150,17 +169,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 + down_samples = 5 # 1s param_name = "LivestreamEncoderBitrate" - def __init__(self, get_stats: Callable[[], dict[str, Any]], params: Params, enabled: bool = True): + def __init__(self, peer_connection: Any, params: Params, enabled: bool = True): super().__init__() - self.get_stats = get_stats + self.pc = peer_connection self.params = params self.level = 2 self._publish(self.bitrates[self.level]) - self.prev_stats: tuple[Any, ...] | None = None + self.prev_lost, self.prev_sent = None, None self.counter = 0 self.up_samples = 5 # 1s self._auto = True @@ -177,7 +196,7 @@ class LivestreamBitrateController(AsyncTaskRunner): if not self._auto: continue - loss_rate = self._sample() + loss_rate = await self._sample() if loss_rate is None: continue if loss_rate >= self.med_level and self.level > 0: @@ -194,18 +213,22 @@ class LivestreamBitrateController(AsyncTaskRunner): self.counter = 0 self._publish(self.bitrates[self.level]) - def _sample(self) -> float | None: - report = next(iter(self.get_stats().values()), None) - if report is None: - return None + 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 - current = (report.ssrc, report.fraction_lost, report.packets_lost, report.highest_seq_no, report.jitter, report.lsr, report.dlsr) - if self.prev_stats == current: + if self.prev_lost is None: + self.prev_lost, self.prev_sent = packets_lost, packets_sent return None - self.prev_stats = current - - loss_rate = report.fraction_lost / 256 - return loss_rate + 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 def _publish(self, bitrate: float): self.params.put(self.param_name, bitrate) @@ -221,41 +244,48 @@ class LivestreamBitrateController(AsyncTaskRunner): class StreamSession: shared_pub_master = DynamicPubMaster([]) - def __init__(self, body: StreamRequestBody): + def __init__(self, body: StreamRequestBody, debug_mode: bool = False): + if debug_mode: + from aiortc.mediastreams import AudioStreamTrack, VideoStreamTrack + from aiortc.contrib.media import MediaBlackhole 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, bind_address=_default_route_ip()) + builder = WebRTCAnswerBuilder(body.sdp) + config = parse_info_from_offer(body.sdp) self.enabled = body.enabled - 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.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.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.get_receiver_report_stats, self.params, self.enabled) + self.bitrate_controller = LivestreamBitrateController(self.stream.peer_connection, 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), 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, + "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, ) def start(self): @@ -280,21 +310,16 @@ class StreamSession: match msg_type: case "livestreamCameraSwitch": - # only needed for 1 track stream - if len(self.video_tracks) == 1: - self.video_tracks[0].switch_camera(payload["data"]["camera"]) + self.video_track.switch_camera(payload["data"]["camera"]) case "livestreamSettings": - if self.bitrate_controller is not None: - self.bitrate_controller.set_quality(payload["data"]["quality"]) + self.bitrate_controller.set_quality(payload["data"]["quality"]) case "livestreamVideoEnable": enabled = payload["data"]["enabled"] self.enabled = enabled - for track in self.video_tracks: - track.enable(enabled) + self.video_track.enable(enabled) if self.outgoing_bridge is not None: self.outgoing_bridge.enable(enabled) - if self.bitrate_controller is not None: - self.bitrate_controller.enable(enabled) + self.bitrate_controller.enable(enabled) if not enabled: self.params.put("LivestreamRequestKeyframe", True) case "clockSync": @@ -303,29 +328,15 @@ class StreamSession: }}) self.stream.get_messaging_channel().send(pong) case "enableTimingSei": - for track in self.video_tracks: - track.timing_sei_enabled = bool(payload["data"]["enabled"]) + if hasattr(self.video_track, 'timing_sei_enabled'): + self.video_track.timing_sei_enabled = bool(payload["data"]["enabled"]) case _: - if msg_type not in self.incoming_bridge_services: + if payload.get("type") not in self.incoming_bridge_services: return - if self.incoming_bridge is not None: - self.incoming_bridge.send(message) + 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) @@ -338,14 +349,15 @@ class StreamSession: channel = self.stream.get_messaging_channel() self.outgoing_bridge.add_channel(channel) self.outgoing_bridge.start() - if self.bitrate_controller is not None: - self.bitrate_controller.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() self.logger.info("Stream session (%s) connected", self.identifier) - if self.is_body: - await self.run_body_session() - else: - await self.run_normal_session() + await self.stream.wait_for_disconnection() self.logger.info("Stream session (%s) ended", self.identifier) except Exception: self.logger.exception("Stream session failure") @@ -358,52 +370,39 @@ class StreamSession: return self._cleanup_done = True self.params.put("LivestreamRequestKeyframe", False) - if self.bitrate_controller is not None: - await self.bitrate_controller.stop() + await self.bitrate_controller.stop() if self.outgoing_bridge is not None: await self.outgoing_bridge.stop() - for track in self.video_tracks: - track.stop() - self.video_tracks.clear() + 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 await self.stream.stop() -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 schedule_teardown(app): + # if nothing connects for 5 seconds, tear down livestreaming processes + h = app.get('teardown') + if h: + h.cancel() def clear(): - if not state.streams: - Params().put_bool("IsLiveStreaming", False) - - state.teardown = asyncio.get_running_loop().call_later(5.0, clear) + if not app['streams']: + Params().put_bool("IsLiveStreaming", False) + app['teardown'] = asyncio.get_running_loop().call_later(5.0, clear) -def _json_response(obj: Any, status: int = 200) -> tuple[int, bytes, str]: - return (status, json.dumps(obj).encode(), "application/json; charset=utf-8") +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 _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: + async with request.app['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 _json_response({"error": "busy", "message": "someone else is connected."}) + return web.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(): @@ -415,15 +414,10 @@ async def handle_get_stream(state: ServerState, raw_body: bytes) -> tuple[int, b await s.stop() stream_dict.pop(sid, None) - session = StreamSession(body) + session = StreamSession(body, debug_mode) stream_dict[session.identifier] = session try: - 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 + answer = await session.get_answer() except Exception: await session.stop() stream_dict.pop(session.identifier, None) @@ -433,189 +427,94 @@ async def handle_get_stream(state: ServerState, raw_body: bytes) -> tuple[int, b def remove_finished_session(_: asyncio.Task) -> None: stream_dict.pop(session.identifier, None) - schedule_teardown(state) - + schedule_teardown(request.app) session.run_task.add_done_callback(remove_finished_session) - return _json_response({"sdp": answer.sdp, "type": answer.type}) + return web.json_response({"sdp": answer.sdp, "type": answer.type}) -async def handle_get_schema(state: ServerState, services_param: str) -> tuple[int, bytes, str]: - services = services_param.split(",") +async def get_schema(request: 'web.Request'): + services = request.query.get("services", "").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 _json_response(schema_dict) + return web.json_response(schema_dict) -async def handle_post_notify(state: ServerState, payload: Any) -> tuple[int, bytes, str]: - for session in list(state.streams.values()): +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()): try: ch = session.stream.get_messaging_channel() ch.send(json.dumps(payload)) except Exception: continue - return _text_response("OK") + return web.Response(status=200, text="OK") -async def on_shutdown(state: ServerState): - for session in list(state.streams.values()): +async def on_shutdown(app: 'web.Application'): + for session in list(app['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() - state.streams.clear() + del app['streams'] -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 +@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 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: +def prewarm_stream_session_imports(debug_mode: bool = False) -> None: + if debug_mode: + from aiortc.mediastreams import VideoStreamTrack + assert VideoStreamTrack from openpilot.system.webrtc.device.video import LiveStreamVideoStreamTrack from teleoprtc.builder import WebRTCAnswerBuilder assert LiveStreamVideoStreamTrack assert WebRTCAnswerBuilder -def webrtcd_thread(host: str, port: int): - logging.basicConfig(level=logging.INFO, handlers=[logging.StreamHandler()]) +def webrtcd_thread(host: str, port: int, debug: bool): + logging.basicConfig(level=logging.CRITICAL, handlers=[logging.StreamHandler()]) prewarm_start = time.monotonic() - prewarm_stream_session_imports() + prewarm_stream_session_imports(debug) prewarm_end = time.monotonic() logging.getLogger("webrtcd").info(f"webrtc prewarm finished in {(prewarm_end - prewarm_start) * 1000} ms") - loop = asyncio.new_event_loop() - asyncio.set_event_loop(loop) - state = ServerState() + app = web.Application(middlewares=[error_middleware]) - server = WebrtcdHTTPServer((host, port), WebrtcdHandler) - server.state = state - server.loop = loop + 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) - # 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() + web.run_app(app, host=host, port=port) 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) + webrtcd_thread(args.host, args.port, args.debug) if __name__=="__main__": diff --git a/teleoprtc_repo/teleoprtc/builder.py b/teleoprtc_repo/teleoprtc/builder.py index cc6565a47..cb9f5d057 100644 --- a/teleoprtc_repo/teleoprtc/builder.py +++ b/teleoprtc_repo/teleoprtc/builder.py @@ -1,7 +1,9 @@ import abc -from typing import Dict, List, Optional +from typing import Dict, List -from teleoprtc.stream import RTCSessionDescription, WebRTCBaseStream, WebRTCOfferStream, WebRTCAnswerStream, ConnectionProvider +import aiortc + +from teleoprtc.stream import WebRTCBaseStream, WebRTCOfferStream, WebRTCAnswerStream, ConnectionProvider from teleoprtc.tracks import TiciVideoStreamTrack, TiciTrackWrapper @@ -12,12 +14,11 @@ class WebRTCStreamBuilder(abc.ABC): class WebRTCOfferBuilder(WebRTCStreamBuilder): - def __init__(self, connection_provider: ConnectionProvider, bind_address: Optional[str] = None): + def __init__(self, connection_provider: ConnectionProvider): self.connection_provider = connection_provider - self.bind_address = bind_address self.requested_camera_types: List[str] = [] self.requested_audio = False - self.audio_tracks: List[object] = [] + self.audio_tracks: List[aiortc.MediaStreamTrack] = [] self.messaging_enabled = False def offer_to_receive_video_stream(self, camera_type: str): @@ -27,7 +28,7 @@ class WebRTCOfferBuilder(WebRTCStreamBuilder): def offer_to_receive_audio_stream(self): self.requested_audio = True - def add_audio_stream(self, track: object): + def add_audio_stream(self, track: aiortc.MediaStreamTrack): assert len(self.audio_tracks) == 0 self.audio_tracks = [track] @@ -42,34 +43,32 @@ 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, bind_address: Optional[str] = None): + def __init__(self, offer_sdp: str): self.offer_sdp = offer_sdp - self.bind_address = bind_address - self.video_tracks: Dict[str, TiciVideoStreamTrack] = {} + self.video_tracks: Dict[str, aiortc.MediaStreamTrack] = dict() self.requested_audio = False - self.audio_tracks: List[object] = [] + self.audio_tracks: List[aiortc.MediaStreamTrack] = [] def offer_to_receive_audio_stream(self): self.requested_audio = True - def add_video_stream(self, camera_type: str, track: object): + def add_video_stream(self, camera_type: str, track: aiortc.MediaStreamTrack): 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: object): + def add_audio_stream(self, track: aiortc.MediaStreamTrack): assert len(self.audio_tracks) == 0 self.audio_tracks = [track] def stream(self) -> WebRTCBaseStream: - description = RTCSessionDescription(sdp=self.offer_sdp, type="offer") + description = aiortc.RTCSessionDescription(sdp=self.offer_sdp, type="offer") return WebRTCAnswerStream( description, consumed_camera_types=[], @@ -77,5 +76,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 deleted file mode 100644 index 2a2cb96e0..000000000 --- a/teleoprtc_repo/teleoprtc/decoder.py +++ /dev/null @@ -1,50 +0,0 @@ -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 191d177cd..537b71242 100644 --- a/teleoprtc_repo/teleoprtc/info.py +++ b/teleoprtc_repo/teleoprtc/info.py @@ -1,6 +1,6 @@ import dataclasses -from libdatachannel import Description +import aiortc @dataclasses.dataclass @@ -15,23 +15,13 @@ def parse_info_from_offer(sdp: str) -> StreamingMediaInfo: """ helper function to parse info about outgoing and incoming streams from an offer sdp """ - desc = Description(sdp, Description.Type.Offer) - n_video = 0 - expected_audio_track = False - incoming_audio_track = False - incoming_datachannel = desc.has_application() + 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 - 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(len(video_tracks), expects_outgoing_audio_track, has_incoming_audio_track, has_incoming_datachannel) - 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 cc4f0092d..208bd6965 100644 --- a/teleoprtc_repo/teleoprtc/stream.py +++ b/teleoprtc_repo/teleoprtc/stream.py @@ -1,29 +1,13 @@ import abc import asyncio -import contextlib import dataclasses import logging -import random -from typing import Any, Awaitable, Callable, Dict, List, Optional, Tuple, Union +from typing import Any, Awaitable, Callable, Dict, List, Optional -from libdatachannel import ( - Configuration, - DataChannel, - Description, - FrameInfo, - H264RtpPacketizer, - IceServer, - NalUnit, - PeerConnection, - PliHandler, - RtcpNackResponder, - RtcpSrReporter, - RtpPacketizationConfig, - Track, -) +import aiortc +from aiortc.contrib.media import MediaRelay -from teleoprtc.decoder import RtcpReceiverReport, _decode_receiver_reports -from teleoprtc.tracks import TiciVideoStreamTrack, parse_video_track_id +from teleoprtc.tracks import parse_video_track_id @dataclasses.dataclass @@ -32,238 +16,138 @@ class StreamingOffer: video: List[str] -@dataclasses.dataclass -class RTCSessionDescription: - sdp: str - type: str - - -ConnectionProvider = Callable[[StreamingOffer], Awaitable[RTCSessionDescription]] -MessageHandler = Callable[[Union[bytes, str]], None] +ConnectionProvider = Callable[[StreamingOffer], Awaitable[aiortc.RTCSessionDescription]] +MessageHandler = Callable[[bytes], Awaitable[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[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) + 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() 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, Any] = {} - self.incoming_audio_tracks: List[Any] = [] - self.outgoing_video_tracks = video_producer_tracks - self.outgoing_audio_tracks = audio_producer_tracks + 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.should_add_data_channel = should_add_data_channel - self.messaging_channel: Optional[DataChannel] = None + self.messaging_channel: Optional[aiortc.RTCDataChannel] = 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_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.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.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 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 + for _ in self.expected_incoming_camera_types: + self.peer_connection.addTransceiver("video", direction="recvonly") if self.expected_incoming_audio: - 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) + self.peer_connection.addTransceiver("audio", direction="recvonly") - 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") + 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 _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 + return target_transceiver - def _add_producer_tracks(self, remote_sdp: Optional[str] = None): - used_mids: set[str] = set() + def _add_producer_tracks(self): for track in self.outgoing_video_tracks: - media, ssrc, payload_type, cname = self._make_video_media(track, remote_sdp or "", used_mids) - rtc_track = self.peer_connection.add_track(media) + 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) + 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") - 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) + self.peer_connection.addTrack(track) - 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)) + def _add_messaging_channel(self, channel: Optional[aiortc.RTCDataChannel] = None): + if not channel: + channel = self.peer_connection.createDataChannel("data", ordered=True) - if self.outgoing_audio_tracks: - raise NotImplementedError("Audio producer tracks are not implemented with libdatachannel") + for handler in self.incoming_message_handlers: + channel.on("message", handler) - def _add_messaging_channel(self, channel: Optional[DataChannel] = None): - if channel is None: - channel = self.peer_connection.create_data_channel("data") + if channel.readyState == "open": + self.messaging_channel_ready_event.set() + else: + channel.on("open", lambda: self.messaging_channel_ready_event.set()) self.messaging_channel = channel - def on_message(message: Union[bytes, str]): - for handler in list(self.incoming_message_handlers): - self._call_soon_threadsafe(handler, message) + 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_open(): - self._set_event(self.messaging_channel_ready_event) + 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_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) + 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) self._on_after_media() - 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: + 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: self._add_messaging_channel(channel) - - 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 + self._on_after_media() def _on_after_media(self): - 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) + if self._number_of_incoming_media == self.expected_number_of_incoming_media: + self.incoming_media_ready_event.set() def _parse_incoming_streams(self, remote_sdp: str): - 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) + 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 def has_incoming_video_track(self, camera_type: str) -> bool: return camera_type in self.incoming_camera_tracks @@ -274,122 +158,65 @@ 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) -> Track: + def get_incoming_video_track(self, camera_type: str, buffered: bool = False) -> aiortc.MediaStreamTrack: 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] - def get_incoming_audio_track(self) -> Track: + 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: 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] - def get_messaging_channel(self) -> DataChannel: + 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: 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) + return self.messaging_channel 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.local_description() is not None and \ - self.peer_connection.remote_description() is not None and \ - self.peer_connection.state() != PeerConnection.State.Closed + self.peer_connection.localDescription is not None and \ + self.peer_connection.remoteDescription is not None and \ + self.peer_connection.connectionState != "closed" @property def is_connected_and_ready(self) -> bool: return self.peer_connection is not None and \ - self.peer_connection.state() == PeerConnection.State.Connected and \ + self.peer_connection.connectionState == "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.state() != PeerConnection.State.Connected: + if self.peer_connection.connectionState != '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): - 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() + await self.peer_connection.close() @abc.abstractmethod - async def start(self) -> RTCSessionDescription: + async def start(self) -> aiortc.RTCSessionDescription: raise NotImplementedError @@ -398,46 +225,76 @@ class WebRTCOfferStream(WebRTCBaseStream): super().__init__(*args, **kwargs) self.session_provider = session_provider - async def start(self) -> RTCSessionDescription: - self._loop = asyncio.get_running_loop() + async def start(self) -> aiortc.RTCSessionDescription: self._add_consumer_transceivers() if self.should_add_data_channel: self._add_messaging_channel() + self._add_producer_tracks() - self.peer_connection.set_local_description(Description.Type.Offer) - await self._wait_for_gathering_complete() - actual_offer = self.peer_connection.local_description() + offer = await self.peer_connection.createOffer() + await self.peer_connection.setLocalDescription(offer) + actual_offer = self.peer_connection.localDescription streaming_offer = StreamingOffer( - sdp=str(actual_offer), + sdp=actual_offer.sdp, video=list(self.expected_incoming_camera_types), ) remote_answer = await self.session_provider(streaming_offer) self._parse_incoming_streams(remote_sdp=remote_answer.sdp) - self.peer_connection.set_remote_description(Description(remote_answer.sdp, Description.Type.Answer)) - self._on_after_media() - actual_answer = self.peer_connection.remote_description() + await self.peer_connection.setRemoteDescription(remote_answer) + actual_answer = self.peer_connection.remoteDescription - return RTCSessionDescription(str(actual_answer), actual_answer.type_string()) + return actual_answer class WebRTCAnswerStream(WebRTCBaseStream): - _retain_messaging_channel_on_close = True - - def __init__(self, session: RTCSessionDescription, *args, **kwargs): + def __init__(self, session: aiortc.RTCSessionDescription, *args, **kwargs): super().__init__(*args, **kwargs) self.session = session - async def start(self) -> RTCSessionDescription: - self._loop = asyncio.get_running_loop() - assert self.peer_connection.remote_description() is None, "Connection already established" + 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) self._parse_incoming_streams(remote_sdp=self.session.sdp) - self.peer_connection.set_remote_description(Description(self.session.sdp, Description.Type.Offer)) - self._add_producer_tracks(self.session.sdp) + await self.peer_connection.setRemoteDescription(self.session) - self.peer_connection.set_local_description(Description.Type.Answer) - await self._wait_for_gathering_complete() - actual_answer = self.peer_connection.local_description() + 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 - return RTCSessionDescription(str(actual_answer), actual_answer.type_string()) diff --git a/teleoprtc_repo/teleoprtc/tracks.py b/teleoprtc_repo/teleoprtc/tracks.py index 12bcc88e7..10b234aa2 100644 --- a/teleoprtc_repo/teleoprtc/tracks.py +++ b/teleoprtc_repo/teleoprtc/tracks.py @@ -1,11 +1,11 @@ -import fractions +import asyncio import logging -import uuid -from typing import Any, Tuple +import time +import fractions +from typing import Any, Optional, Tuple - -VIDEO_CLOCK_RATE = 90000 -VIDEO_TIME_BASE = fractions.Fraction(1, VIDEO_CLOCK_RATE) +import aiortc +from aiortc.mediastreams import VIDEO_CLOCK_RATE, VIDEO_TIME_BASE def video_track_id(camera_type: str, track_id: str) -> str: @@ -21,51 +21,57 @@ def parse_video_track_id(track_id: str) -> Tuple[str, str]: return camera_type, track_id -class TiciVideoStreamTrack: +class TiciVideoStreamTrack(aiortc.MediaStreamTrack): """ - 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"] - self._id: str = video_track_id(camera_type, str(uuid.uuid4())) + 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._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 recv(self): - raise NotImplementedError() + async def next_pts(self, current_pts) -> float: + pts: float = current_pts + self._dt * self._clock_rate - def request_keyframe(self) -> None: - pass + 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 -class TiciTrackWrapper(TiciVideoStreamTrack): +class TiciTrackWrapper(aiortc.MediaStreamTrack): """ - Associates a generic video track with camera_type. + Associates video track with camera_type """ - def __init__(self, camera_type: str, track: Any): + def __init__(self, camera_type: str, track: aiortc.MediaStreamTrack): assert track.kind == "video" - super().__init__(camera_type, getattr(track, "_dt", 0.05)) + assert not isinstance(track, TiciVideoStreamTrack) + super().__init__() 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 7d91188ad..e13580b5d 100755 --- a/tools/agnos/flash_desktop_system_to_comma.sh +++ b/tools/agnos/flash_desktop_system_to_comma.sh @@ -2,45 +2,12 @@ set -euo pipefail HOST="${1:-comma@192.168.3.110}" -IMAGE="${2:-/Users/dominickthompson/Desktop/system17.img.xz}" -METADATA="${3:-${IMAGE}.metadata.json}" +IMAGE="${2:-/Users/dominickthompson/Desktop/system8.img.xz}" SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)" REPO_ROOT="$(cd -- "${SCRIPT_DIR}/../.." && pwd)" -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 -} +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) SESSION="local_agnos_flash" REMOTE_DIR="/data/local_agnos_flash" @@ -48,49 +15,40 @@ 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( - needle, - '''class _UnusedCasync: - ChunkReader = object - ChunkDict = object - + "import openpilot.system.updated.casync.casync as casync", + """class _UnusedCasync: 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 -PYTHON_BIN="/usr/local/venv/bin/python3" -[[ -x "$PYTHON_BIN" ]] || { echo "[ERROR] managed Python is unavailable" >&2; exit 1; } +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 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" - "$REMOTE_MANIFEST" <<'PY' -import json + if "$PYTHON_BIN" - "${PORT}" "${IMAGE_NAME}" <<'PY' import sys import urllib.request -from pathlib import Path -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) +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) PY then http_ready=1 @@ -150,23 +115,24 @@ PY sleep 0.25 done -[[ "${http_ready:-0}" == "1" ]] || { - echo "[ERROR] local image server did not become ready" >&2 +if [[ "$http_ready" != "1" ]]; then + echo "[ERROR] Local image HTTP server did not become ready" >&2 cat "${REMOTE_DIR}/http.log" >&2 || true exit 1 -} +fi -echo "[FLASH] Writing and verifying the candidate in the inactive system slot" +echo "[FLASH] Flashing local system image to inactive AGNOS slot" PYTHONPATH="$(dirname "$REMOTE_AGNOS")" "$PYTHON_BIN" "$REMOTE_AGNOS" --swap "$REMOTE_MANIFEST" -echo "[DONE] Candidate written, verified, and selected" +echo "[DONE] AGNOS flashed and slot swapped" +echo "[REBOOT] Rebooting now" 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' 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' IMAGE_NAME='$IMAGE_NAME' EXPECTED_VERSION='$EXPECTED_VERSION' bash '$REMOTE_RUNNER'\"" echo "Started remote tmux session: $SESSION" -echo "After reboot, run tools/agnos/validate_agnos_runtime.sh $EXPECTED_VERSION on the device." +echo "Watch it with: ssh $HOST 'tmux attach -t $SESSION'" diff --git a/tools/agnos/patch_system_reset_image.py b/tools/agnos/patch_system_reset_image.py old mode 100755 new mode 100644 index 1bd27251c..9b9f1c595 --- a/tools/agnos/patch_system_reset_image.py +++ b/tools/agnos/patch_system_reset_image.py @@ -1,149 +1,50 @@ #!/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 struct import subprocess +import struct import tempfile import urllib.request import zipfile +from io import BytesIO from pathlib import Path -VERSION_PATH_IN_IMAGE = "/VERSION" -SITE_PACKAGES_PATH_IN_IMAGE = "/usr/local/venv/lib/python3.12/site-packages" -LEGACY_RUNTIME_LIBRARY_DIR = "/usr/local/lib" +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" -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", -} - +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" ANDROID_SPARSE_MAGIC = 0xED26FF3A CHUNK_TYPE_RAW = 0xCAC1 CHUNK_TYPE_FILL = 0xCAC2 @@ -151,878 +52,1312 @@ 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: - 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") + 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() def find_debugfs() -> str: - return find_tool("debugfs", ("/opt/homebrew/opt/e2fsprogs/sbin/debugfs",)) - - -def find_e2fsck() -> str: - return find_tool("e2fsck", ("/opt/homebrew/opt/e2fsprogs/sbin/e2fsck",)) - - -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: - digest = hashlib.sha256() - with path.open("rb") as stream: - while chunk := stream.read(8 * 1024 * 1024): - digest.update(chunk) - return digest.hexdigest() - - -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 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 install_bundled_installer(owner: str, branch: str, installer_url: str) -> None: - import tempfile - - 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, - ) - - source = replace_exactly( - source, - ''' # 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()''', - ''' # 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()''', - ) - - 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 - - import tempfile -''', - ) - - 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 - - -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) + 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.") + + +def load_manifest(path: Path) -> list[dict]: + return json.loads(path.read_text()) def get_system_entry(manifest: list[dict]) -> dict: - return next(entry for entry in manifest if entry.get("name") == "system") + for e in manifest: + if e.get("name") == "system": + return e + raise RuntimeError("No system entry found in manifest") -def update_manifest_system_entry(manifest: list[dict], new_url: str, raw_hash: str, size: int) -> list[dict]: +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 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) + 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 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 update_manifest_system_entry(manifest: list[dict], new_url: str, new_hash_raw: 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) + 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") + + 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 + + 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, + } + return updated +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 + + 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) + + 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." + ) + + 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 main() -> int: args = parse_args() - target_version = validate_target_version(args.set_version) - debugfs, e2fsck = find_debugfs(), find_e2fsck() + debugfs = find_debugfs() + + manifest_path = Path(args.manifest).resolve() + manifest = load_manifest(manifest_path) + system_entry = get_system_entry(manifest) + work_dir = Path(args.work_dir).resolve() work_dir.mkdir(parents=True, exist_ok=True) if args.source_image: - source = Path(args.source_image).resolve() + downloaded_img = Path(args.source_image).resolve() + if not downloaded_img.is_file(): + raise RuntimeError(f"Source image not found: {downloaded_img}") else: - 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}") + 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) - 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 + 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", } - if upstream_differences: - raise RuntimeError(f"Source image is not exact upstream AGNOS: {json.dumps(upstream_differences, sort_keys=True)}") + 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) - 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_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" - 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) + 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) - 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("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) - 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) + 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) - 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") + 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) - 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") + 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("") if args.new_url: - 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") + 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)) + return 0 diff --git a/tools/agnos/test_patch_system_reset_image.py b/tools/agnos/test_patch_system_reset_image.py index edec1d23c..c2742978a 100644 --- a/tools/agnos/test_patch_system_reset_image.py +++ b/tools/agnos/test_patch_system_reset_image.py @@ -1,238 +1,106 @@ -import importlib.util -import json -import os -import zipfile from pathlib import Path +import runpy import pytest - -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 +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, +) -patch_image = _load_patch_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 +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 -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} +exec /data/continue.sh ''' -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) +def test_patch_comma_sh_display_wait_uses_available_display_service(): + patched = patch_comma_sh_display_wait(ORIGINAL_DISPLAY_WAIT) - 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 + 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 -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 +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_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("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_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) +@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() diff --git a/tools/agnos/validate_agnos_runtime.sh b/tools/agnos/validate_agnos_runtime.sh deleted file mode 100755 index 992db70c2..000000000 --- a/tools/agnos/validate_agnos_runtime.sh +++ /dev/null @@ -1,90 +0,0 @@ -#!/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'