Compare commits

..

1 Commits

Author SHA1 Message Date
firestar5683 b56eca9428 DIE AGNOS DIE 2026-08-24 12:36:15 -05:00
24 changed files with 1658 additions and 2010 deletions
+32 -2
View File
@@ -1,4 +1,5 @@
import os
import importlib
import shutil
import subprocess
import sys
@@ -128,6 +129,23 @@ 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 []
# Homebrew llvm can shadow Apple clang and break macOS SDK header resolution.
# Use the system toolchain explicitly on macOS for reliable local builds.
cc = '/usr/bin/clang' if arch == "Darwin" else 'clang'
@@ -270,6 +288,8 @@ env = Environment(
] + cflags + ccflags,
CPPPATH=cpppath + [
capnproto_include_dirs,
ffmpeg_include_dirs,
"#",
"#third_party/acados/include",
"#third_party/acados/include/blasfeo/include",
@@ -288,11 +308,13 @@ env = Environment(
RANLIB=ranlib,
LINKFLAGS=ldflags,
RPATH=rpath,
RPATH=rpath + ffmpeg_lib_dirs,
CFLAGS=["-std=gnu11"] + cflags,
CXXFLAGS=["-std=c++1z"] + cxxflags,
LIBPATH=libpath + [
capnproto_lib_dirs,
ffmpeg_lib_dirs,
"#msgq_repo",
"#third_party",
"#selfdrive/pandad",
@@ -371,7 +393,15 @@ SConscript(['opendbc_repo/SConscript'], exports={'env': env_swaglog})
SConscript(['cereal/SConscript'])
Import('socketmaster', 'msgq')
messaging = [socketmaster, msgq, 'capnp', 'kj',]
if capnproto is not None:
messaging = [
socketmaster,
msgq,
File(os.path.join(capnproto.LIB_DIR, "libcapnp.a")),
File(os.path.join(capnproto.LIB_DIR, "libkj.a")),
]
else:
messaging = [socketmaster, msgq, 'capnp', 'kj']
Export('messaging')
+1 -1
View File
@@ -21,7 +21,7 @@ fi
export QCOM_PRIORITY=12
if [ -z "$AGNOS_VERSION" ]; then
export AGNOS_VERSION="19.6.2"
export AGNOS_VERSION="19.6.6"
fi
if [ -z "$AGNOS_ACCEPTED_VERSIONS" ]; then
+5
View File
@@ -29,6 +29,11 @@ dependencies = [
"setuptools",
"numpy >=2.0",
# AGNOS 19.6 native build dependencies
"comma-deps-capnproto; python_version >= '3.12'",
"comma-deps-ffmpeg; python_version >= '3.12'",
"libdatachannel-py>=2026.1.0.dev2; python_version >= '3.12'",
# body / webrtcd
"aiohttp",
"aiortc",
+24 -6
View File
@@ -1,15 +1,29 @@
import importlib.util
import os
from pathlib import Path
Import('env', 'arch', 'common')
if GetOption('extras') and arch != "Darwin":
raylib_env = env.Clone()
raylib_env['LIBPATH'] += [f'#third_party/raylib/{arch}/']
raylib_env['LINKFLAGS'].append('-Wl,-strip-debug')
raylib_libs = common + ["raylib"]
if arch == "larch64":
raylib_libs += ["GLESv2", "wayland-client", "wayland-egl", "EGL"]
tici_sysroot = os.environ.get("SP_TICI_SYSROOT", "").strip().rstrip("/")
if tici_sysroot:
raylib_dir = Path(tici_sysroot) / "usr/local/venv/lib/python3.12/site-packages/raylib/install"
else:
raylib_spec = importlib.util.find_spec("raylib")
if raylib_spec is None or raylib_spec.submodule_search_locations is None:
raise RuntimeError("The managed raylib package is required to build the comma installer")
raylib_dir = Path(raylib_spec.submodule_search_locations[0]) / "install"
raylib_env['CPPPATH'] += [str(raylib_dir / "include")]
raylib_env['LIBPATH'] += [str(raylib_dir / "lib")]
raylib_libs = common + ["raylib_comma", "GLESv2", "EGL", "gbm", "drm"]
else:
raylib_libs += ["GL"]
raylib_env['LIBPATH'] += [f'#third_party/raylib/{arch}/']
raylib_libs = common + ["raylib", "GL"]
release = "release3"
installers = [
@@ -23,11 +37,15 @@ if GetOption('extras') and arch != "Darwin":
"ld -r -b binary -o $TARGET $SOURCE")
inter = raylib_env.Command("installer/inter_ttf.o", "installer/inter-ascii.ttf",
"ld -r -b binary -o $TARGET $SOURCE")
inter_bold = raylib_env.Command("installer/inter_bold.o", "../assets/fonts/Inter-Bold.ttf",
"ld -r -b binary -o $TARGET $SOURCE")
inter_light = raylib_env.Command("installer/inter_light.o", "../assets/fonts/Inter-Light.ttf",
"ld -r -b binary -o $TARGET $SOURCE")
for name, branch in installers:
defines = {'BRANCH': f"'\"{branch}\"'"}
if "internal" in name:
defines['INTERNAL'] = "1"
obj = raylib_env.Object(f"installer/installers/installer_{name}.o", ["installer/installer.cc"], CPPDEFINES=defines)
installer = raylib_env.Program(f"installer/installers/installer_{name}", [obj, cont, inter], LIBS=raylib_libs)
assert installer[0].get_size() < 1900*1e3, installer[0].get_size()
installer = raylib_env.Program(f"installer/installers/installer_{name}", [obj, cont, inter, inter_bold, inter_light], LIBS=raylib_libs)
assert installer[0].get_size() < 2500*1e3, installer[0].get_size()
+1 -1
View File
@@ -6,7 +6,7 @@
#include "common/swaglog.h"
#include "common/util.h"
#include "system/hardware/hw.h"
#include "third_party/raylib/include/raylib.h"
#include "raylib.h"
int freshClone();
int cachedFetch(const std::string &cache);
@@ -697,105 +697,73 @@ class FavoriteRadialMenu:
lines[-1] = FavoriteRadialMenu._append_ellipsis(font, lines[-1], font_size, max_width)
return lines
@staticmethod
def _sample_aether_color(t: float, alpha: int) -> rl.Color:
"""Sample from the tri-stop Aether gradient (#C49EFF -> #AC7DFF -> #58329E)."""
if t <= 0.5:
w = t * 2.0
r = int(196 - 24 * w)
g = int(158 - 33 * w)
b = 255
else:
w = (t - 0.5) * 2.0
r = int(172 - 84 * w)
g = int(125 - 75 * w)
b = int(255 - 97 * w)
return rl.Color(r, g, b, alpha)
def _draw_corner_hint(self) -> None:
scale = self._scale_for(self._rect)
x0 = self._rect.x
y0 = self._rect.y + self._rect.height
is_pressed = self._corner_press is not None
size = 150.0 * scale
steps = 48
# 1. Precompute edge vertices to minimize per-frame allocations
v_origin = rl.Vector2(x0, y0)
inv_steps = 1.0 / steps
pts_top = [rl.Vector2(x0, y0 - size * (k * inv_steps)) for k in range(steps + 1)]
pts_right = [rl.Vector2(x0 + size * (k * inv_steps), y0) for k in range(steps + 1)]
# 2. Pass 0 (Smoked Obsidian Base) & Pass 1 (Tri-Stop Aether Gradient Mesh)
# Borderless design: smooth monotonic decay into exact 0 alpha at hypotenuse
base_max_alpha = 200 if is_pressed else 160
purple_max_alpha = 145 if is_pressed else 105
size = 120.0 * scale
purple = self._PURPLE
steps = 6
max_alpha = 135 if is_pressed else 95
for i in range(steps):
t_mid = (i + 0.5) * inv_steps
v_ta, v_tb = pts_top[i], pts_top[i + 1]
v_ra, v_rb = pts_right[i], pts_right[i + 1]
t_a = i / float(steps)
t_b = (i + 1) / float(steps)
t_mid = (t_a + t_b) * 0.5
base_a = int(base_max_alpha * ((1.0 - t_mid) ** 1.40))
purple_a = int(purple_max_alpha * ((1.0 - t_mid) ** 1.75))
# Non-linear quadratic fade: dissolves cleanly into transparent road video
alpha = int(max_alpha * ((1.0 - t_mid) ** 1.7))
if alpha <= 0:
continue
for col in (rl.Color(8, 6, 18, base_a) if base_a > 0 else None,
self._sample_aether_color(t_mid, purple_a) if purple_a > 0 else None):
if col is None:
continue
if i == 0:
rl.draw_triangle(v_origin, v_rb, v_tb, col)
else:
rl.draw_triangle(v_ta, v_ra, v_tb, col)
rl.draw_triangle(v_tb, v_ra, v_rb, col)
slice_col = rl.Color(*purple, alpha)
v_a_top = rl.Vector2(x0, y0 - size * t_a)
v_a_right = rl.Vector2(x0 + size * t_a, y0)
v_b_top = rl.Vector2(x0, y0 - size * t_b)
v_b_right = rl.Vector2(x0 + size * t_b, y0)
# 3. Ultra-Polished Frosted-Glass Vector Arrow (Nestled deep in purple corner)
cx = x0 + 34.0 * scale
cy = y0 - 34.0 * scale
if i == 0:
rl.draw_triangle(rl.Vector2(x0, y0), v_b_right, v_b_top, slice_col)
else:
rl.draw_triangle(v_a_top, v_a_right, v_b_top, slice_col)
rl.draw_triangle(v_b_top, v_a_right, v_b_right, slice_col)
tip = rl.Vector2(cx + 15.0 * scale, cy - 15.0 * scale)
tail = rl.Vector2(cx - 15.0 * scale, cy + 15.0 * scale)
wing1 = rl.Vector2(tip.x - 14.0 * scale, tip.y + 1.2 * scale)
wing2 = rl.Vector2(tip.x - 1.2 * scale, tip.y + 14.0 * scale)
# 2. Elegant, optically centered vector arrow
cx = x0 + size * 0.32
cy = y0 - size * 0.32
line_w = 4.6 * scale
tip = rl.Vector2(cx + 8.0 * scale, cy - 8.0 * scale)
tail = rl.Vector2(cx - 7.0 * scale, cy + 7.0 * scale)
wing1 = rl.Vector2(tip.x - 9.0 * scale, tip.y + 0.5 * scale)
wing2 = rl.Vector2(tip.x - 0.5 * scale, tip.y + 9.0 * scale)
# Tier 1: Deep Subsurface Ambient Occlusion Shadow
s_off = 1.6 * scale
s_tip = rl.Vector2(tip.x + s_off, tip.y + s_off)
s_tail = rl.Vector2(tail.x + s_off, tail.y + s_off)
s_w1 = rl.Vector2(wing1.x + s_off, wing1.y + s_off)
s_w2 = rl.Vector2(wing2.x + s_off, wing2.y + s_off)
shadow_w = line_w + 2.0 * scale
shadow_col = rl.Color(8, 6, 16, 130 if is_pressed else 105)
line_w = 3.0 * scale
arrow_col = rl.Color(255, 255, 255, 235 if is_pressed else 195)
shadow_col = rl.Color(16, 10, 28, 120)
shadow_offset = 1.2 * scale
# Tier 2: Aether Violet Refractive Halo / Glass Bloom
halo_w = line_w + 3.2 * scale
halo_col = rl.Color(185, 145, 255, 75 if is_pressed else 55)
# Soft ambient drop shadow behind arrow
s_tip = rl.Vector2(tip.x + shadow_offset, tip.y + shadow_offset)
s_tail = rl.Vector2(tail.x + shadow_offset, tail.y + shadow_offset)
s_w1 = rl.Vector2(wing1.x + shadow_offset, wing1.y + shadow_offset)
s_w2 = rl.Vector2(wing2.x + shadow_offset, wing2.y + shadow_offset)
rl.draw_line_ex(s_tail, s_tip, line_w + 1.0 * scale, shadow_col)
rl.draw_line_ex(s_tip, s_w1, line_w + 1.0 * scale, shadow_col)
rl.draw_line_ex(s_tip, s_w2, line_w + 1.0 * scale, shadow_col)
# Tier 3: Radiant High-Luminance Frost White Body
arrow_col = rl.Color(255, 255, 255, 245 if is_pressed else 225)
# Crisp arrow strokes
rl.draw_line_ex(tail, tip, line_w, arrow_col)
rl.draw_line_ex(tip, wing1, line_w, arrow_col)
rl.draw_line_ex(tip, wing2, line_w, arrow_col)
# Tier 4: Specular Spine Highlight
spec_w = 2.0 * scale
spec_col = rl.Color(255, 255, 255, 255 if is_pressed else 240)
# Render multi-pass optical stack
layers = (
(s_tail, s_tip, s_w1, s_w2, shadow_w, shadow_col),
(tail, tip, wing1, wing2, halo_w, halo_col),
(tail, tip, wing1, wing2, line_w, arrow_col),
(tail, tip, wing1, wing2, spec_w, spec_col),
)
for p_tail, p_tip, p_w1, p_w2, width, col in layers:
rl.draw_line_ex(p_tail, p_tip, width, col)
rl.draw_line_ex(p_tip, p_w1, width, col)
rl.draw_line_ex(p_tip, p_w2, width, col)
r_cap = width * 0.5
for pt in (p_tail, p_tip, p_w1, p_w2):
rl.draw_circle_v(pt, r_cap, col)
# Smooth rounded joints
r_cap = line_w * 0.5
rl.draw_circle_v(tip, r_cap, arrow_col)
rl.draw_circle_v(tail, r_cap, arrow_col)
rl.draw_circle_v(wing1, r_cap, arrow_col)
rl.draw_circle_v(wing2, r_cap, arrow_col)
def _draw_radial_menu(self) -> None:
scale = self._scale_for(self._rect)
+4 -4
View File
@@ -67,13 +67,13 @@
},
{
"name": "system",
"url": "https://www.dropbox.com/scl/fi/pewhzpqzi3aewuiaffc6m/system10.img.xz?rlkey=olzrzulhs93zzghnjrskmdwxt&st=exnfk2oz&dl=1",
"hash": "ab395d4c963a908ab86709f1a6580a62dd24cb34ee71cf5fd5ec29d7d48d0e10",
"hash_raw": "ab395d4c963a908ab86709f1a6580a62dd24cb34ee71cf5fd5ec29d7d48d0e10",
"url": "https://www.dropbox.com/scl/fi/7b48biwpb0hhjh33b30xk/system14.img.xz?rlkey=vbk2m8otxaeheas7uw0garc2i&st=07v7m6ns&dl=1",
"hash": "adcaa5274fd45c486364ad0ae557208087a7499a4b31a949c6ec2447a74943ec",
"hash_raw": "adcaa5274fd45c486364ad0ae557208087a7499a4b31a949c6ec2447a74943ec",
"size": 4718592000,
"sparse": false,
"full_check": false,
"has_ab": true,
"ondevice_hash": "ab395d4c963a908ab86709f1a6580a62dd24cb34ee71cf5fd5ec29d7d48d0e10"
"ondevice_hash": "adcaa5274fd45c486364ad0ae557208087a7499a4b31a949c6ec2447a74943ec"
}
]
-5
View File
@@ -1094,11 +1094,6 @@ def manager_init() -> None:
device=HARDWARE.get_device_type())
last_timing = _log_boot_timing("manager_init", "logging_ready", manager_init_start, last_timing)
# preimport all processes
for p in managed_processes.values():
p.prepare()
last_timing = _log_boot_timing("manager_init", "preimport_processes", manager_init_start, last_timing)
# StarPilot variables
install_starpilot(build_metadata, params)
last_timing = _log_boot_timing("manager_init", "install_starpilot", manager_init_start, last_timing)
+1 -13
View File
@@ -634,19 +634,7 @@ class PythonProcess(ManagerProcess):
self.launcher = launcher
def prepare(self) -> None:
if self.enabled:
cloudlog.info(f"preimporting {self.module}")
start = time.monotonic()
try:
importlib.import_module(self.module)
finally:
line = f"SP_BOOT_TIMING preimport {self.name} module={self.module} +{time.monotonic() - start:.3f}s"
try:
with open(os.environ.get("SP_BOOT_TIMING_LOG", "/tmp/starpilot_boot_timing.log"), "a") as f:
f.write(line + "\n")
except OSError:
pass
cloudlog.warning(line)
pass
def start(self) -> None:
# In case we only tried a non blocking stop we need to stop it before restarting
+20 -15
View File
@@ -1,17 +1,19 @@
import asyncio
from dataclasses import dataclass
import struct
import time
import av
from teleoprtc.tracks import TiciVideoStreamTrack
from aiortc.mediastreams import MediaStreamError
from cereal import messaging
from openpilot.common.params import Params
from openpilot.common.realtime import DT_MDL
from openpilot.common.params import Params
# v4l2 buffer flag marking an encoded keyframe (linux/videodev2.h)
V4L2_BUF_FLAG_KEYFRAME = 0x8
# arbitrary 16-byte UUID identifying openpilot frame-timing SEI messages
TIMING_SEI_UUID = bytes([
0xa5, 0xe0, 0xc4, 0xa4, 0x5b, 0x6e, 0x4e, 0x1e,
0x9c, 0x7e, 0x12, 0x34, 0x56, 0x78, 0x9a, 0xbc,
@@ -19,6 +21,15 @@ TIMING_SEI_UUID = bytes([
_SEI_PREFIX = b'\x00\x00\x00\x01\x06\x05\x30' + TIMING_SEI_UUID
@dataclass(frozen=True)
class EncodedVideoFrame:
data: bytes
pts: int
def __bytes__(self) -> bytes:
return self.data
class LiveStreamVideoStreamTrack(TiciVideoStreamTrack):
camera_to_sock_mapping = {
"driver": "livestreamDriverEncodeData",
@@ -52,6 +63,9 @@ class LiveStreamVideoStreamTrack(TiciVideoStreamTrack):
if not enabled:
self._seen_keyframe = False
def request_keyframe(self) -> None:
self.params.put("LivestreamRequestKeyframe", True, block=False)
def _build_frame_data(self, msg) -> bytes:
encode_data = getattr(msg, msg.which())
if not self.timing_sei_enabled:
@@ -68,9 +82,7 @@ class LiveStreamVideoStreamTrack(TiciVideoStreamTrack):
async def recv(self):
while True:
if self.readyState != "live":
raise MediaStreamError
# while video is disabled, pause here without returning
if not self.video_enabled:
await asyncio.sleep(0.005)
continue
@@ -79,18 +91,11 @@ class LiveStreamVideoStreamTrack(TiciVideoStreamTrack):
if msg is not None:
if not self._seen_keyframe and (getattr(msg, msg.which()).idx.flags & V4L2_BUF_FLAG_KEYFRAME):
self._seen_keyframe = True
self.params.put("LivestreamRequestKeyframe", False)
self.params.put("LivestreamRequestKeyframe", False, block=False)
break
await asyncio.sleep(0.005)
packet = av.Packet(self._build_frame_data(msg))
packet.time_base = self._time_base
self._pts = ((time.monotonic_ns() - self._t0_ns) * self._clock_rate) // 1_000_000_000
packet.pts = self._pts
self.log_debug("track sending frame %d", self._pts)
return packet
def codec_preference(self) -> str | None:
return "H264"
return EncodedVideoFrame(self._build_frame_data(msg), self._pts)
+20 -44
View File
@@ -1,20 +1,17 @@
import asyncio
import json
import time
# for aiortc and its dependencies
import warnings
warnings.filterwarnings("ignore", category=DeprecationWarning)
warnings.filterwarnings("ignore", category=RuntimeWarning) # TODO: remove this when google-crc32c publish a python3.12 wheel
from aiortc import RTCDataChannel
from aiortc.mediastreams import VIDEO_CLOCK_RATE, VIDEO_TIME_BASE
import capnp
import pyaudio
import pytest
pytest.importorskip("libdatachannel", reason="the upstream WebRTC backend requires Python 3.12")
from cereal import messaging, log
from teleoprtc.tracks import VIDEO_CLOCK_RATE
from openpilot.system.webrtc.webrtcd import CerealOutgoingMessageProxy, CerealIncomingMessageProxy
from openpilot.system.webrtc.device.video import LiveStreamVideoStreamTrack
from openpilot.system.webrtc.device.audio import AudioInputStreamTrack
class TestStreamSession:
@@ -33,40 +30,37 @@ class TestStreamSession:
expected_dict = {"type": "customReservedRawData0", "logMonoTime": 123, "valid": True, "data": "test"}
expected_json = json.dumps(expected_dict).encode()
channel = mocker.Mock(spec=RTCDataChannel)
mocked_submaster = messaging.SubMaster(["customReservedRawData0"])
def mocked_update(t):
mocked_submaster.update_msgs(0, [test_msg])
channel = mocker.Mock()
channel.is_open.return_value = True
proxy = CerealOutgoingMessageProxy(["customReservedRawData0"])
def mocked_update(_):
proxy.sm.update_msgs(0, [test_msg])
mocker.patch.object(messaging.SubMaster, "update", side_effect=mocked_update)
proxy = CerealOutgoingMessageProxy(["customReservedRawData0"])
proxy.sm = mocked_submaster
proxy.add_channel(channel)
proxy.update()
channel.send.assert_called_once_with(expected_json)
def test_incoming_proxy(self, mocker):
tested_msgs = [
{"type": "customReservedRawData0", "data": "test"}, # primitive
{"type": "can", "data": [{"address": 0, "dat": "", "src": 0}]}, # list
{"type": "testJoystick", "data": {"axes": [0, 0], "buttons": [False]}}, # dict
{"type": "customReservedRawData0", "data": "test"},
{"type": "can", "data": [{"address": 0, "dat": "", "src": 0}]},
{"type": "testJoystick", "data": {"axes": [0, 0], "buttons": [False]}},
]
mocked_pubmaster = mocker.MagicMock(spec=messaging.PubMaster)
proxy = CerealIncomingMessageProxy(mocked_pubmaster)
for msg in tested_msgs:
proxy.send(json.dumps(msg).encode())
mocked_pubmaster.send.assert_called_once()
mt, md = mocked_pubmaster.send.call_args.args
assert mt == msg["type"]
assert isinstance(md, capnp._DynamicStructBuilder)
assert hasattr(md, msg["type"])
msg_type, message = mocked_pubmaster.send.call_args.args
assert msg_type == msg["type"]
assert isinstance(message, capnp._DynamicStructBuilder)
assert hasattr(message, msg_type)
mocked_pubmaster.reset_mock()
def test_livestream_track(self, mocker):
@@ -78,29 +72,11 @@ class TestStreamSession:
track = LiveStreamVideoStreamTrack("driver")
assert track.id.startswith("driver")
assert track.codec_preference() == "H264"
for i in range(5):
packet = self.loop.run_until_complete(track.recv())
assert packet.time_base == VIDEO_TIME_BASE
if i == 0:
start_ns = time.monotonic_ns()
start_pts = packet.pts
assert abs(i + packet.pts - (start_pts + (((time.monotonic_ns() - start_ns) * VIDEO_CLOCK_RATE) // 1_000_000_000))) < 450 #5ms
assert packet.size == 0
def test_input_audio_track(self, mocker):
packet_time, rate = 0.02, 16000
sample_count = int(packet_time * rate)
mocked_stream = mocker.MagicMock(spec=pyaudio.Stream)
mocked_stream.read.return_value = b"\x00" * 2 * sample_count
config = {"open.side_effect": lambda *args, **kwargs: mocked_stream}
mocker.patch("pyaudio.PyAudio", spec=True, **config)
track = AudioInputStreamTrack(audio_format=pyaudio.paInt16, packet_time=packet_time, rate=rate)
for i in range(5):
frame = self.loop.run_until_complete(track.recv())
assert frame.rate == rate
assert frame.samples == sample_count
assert frame.pts == i * sample_count
assert abs(i + packet.pts - (start_pts + (((time.monotonic_ns() - start_ns) * VIDEO_CLOCK_RATE) // 1_000_000_000))) < 450
assert bytes(packet) == b""
+28 -51
View File
@@ -1,65 +1,42 @@
import pytest
import asyncio
import json
# for aiortc and its dependencies
import warnings
warnings.filterwarnings("ignore", category=DeprecationWarning)
warnings.filterwarnings("ignore", category=RuntimeWarning) # TODO: remove this when google-crc32c publish a python3.12 wheel
from openpilot.system.webrtc.webrtcd import get_stream
import pytest
import aiortc
from teleoprtc import WebRTCOfferBuilder
from parameterized import parameterized_class
pytest.importorskip("libdatachannel", reason="the upstream WebRTC backend requires Python 3.12")
from openpilot.system.webrtc.webrtcd import ServerState, handle_get_schema, handle_post_notify, on_shutdown
@parameterized_class(("in_services", "out_services"), [
(["testJoystick"], ["carState"]),
([], ["carState"]),
(["testJoystick"], []),
([], []),
])
@pytest.mark.asyncio
class TestWebrtcdProc:
async def assertCompletesWithTimeout(self, awaitable, timeout=1):
try:
async with asyncio.timeout(timeout):
await awaitable
except TimeoutError:
pytest.fail("Timeout while waiting for awaitable to complete")
async def test_get_schema():
status, body, content_type = await handle_get_schema(ServerState(), "carState")
async def test_webrtcd(self, mocker):
mock_request = mocker.MagicMock()
async def connect(offer):
body = {'sdp': offer.sdp, 'init_camera': offer.video[0], 'enabled': True,
'bridge_services_in': self.in_services, 'bridge_services_out': self.out_services}
mock_request.json.side_effect = mocker.AsyncMock(return_value=body)
response = await get_stream(mock_request)
response_json = json.loads(response.text)
return aiortc.RTCSessionDescription(**response_json)
assert status == 200
assert content_type.startswith("application/json")
assert "carState" in json.loads(body)
builder = WebRTCOfferBuilder(connect)
builder.offer_to_receive_video_stream("road")
builder.offer_to_receive_audio_stream()
if len(self.in_services) > 0 or len(self.out_services) > 0:
builder.add_messaging()
stream = builder.stream()
@pytest.mark.asyncio
async def test_get_schema_rejects_unknown_service():
with pytest.raises(AssertionError, match="Invalid service name"):
await handle_get_schema(ServerState(), "notARealService")
await self.assertCompletesWithTimeout(stream.start())
await self.assertCompletesWithTimeout(stream.wait_for_connection())
assert stream.has_incoming_video_track("road")
assert stream.has_incoming_audio_track()
assert stream.has_messaging_channel() == (len(self.in_services) > 0 or len(self.out_services) > 0)
@pytest.mark.asyncio
async def test_notify_and_shutdown_active_stream(mocker):
state = ServerState()
session = mocker.MagicMock()
session.stop = mocker.AsyncMock()
state.streams["test"] = session
video_track, audio_track = stream.get_incoming_video_track("road"), stream.get_incoming_audio_track()
await self.assertCompletesWithTimeout(video_track.recv())
await self.assertCompletesWithTimeout(audio_track.recv())
status, body, content_type = await handle_post_notify(state, {"type": "ping"})
await self.assertCompletesWithTimeout(stream.stop())
assert (status, body) == (200, b"OK")
assert content_type.startswith("text/plain")
channel = session.stream.get_messaging_channel.return_value
channel.send.assert_called_once_with(json.dumps({"type": "ping"}))
# cleanup, very implementation specific, test may break if it changes
assert mock_request.app["streams"].__setitem__.called, "Implementation changed, please update this test"
_, session = mock_request.app["streams"].__setitem__.call_args.args
await self.assertCompletesWithTimeout(session.post_run_cleanup())
await on_shutdown(state)
session.stop.assert_awaited_once()
assert state.streams == {}
+250 -149
View File
@@ -1,33 +1,29 @@
#!/usr/bin/env python3
from abc import abstractmethod
from collections.abc import Callable
import os
import socket
import time
import capnp
import argparse
import asyncio
import contextlib
import json
import uuid
import logging
from typing import Any, TYPE_CHECKING
# aiortc and its dependencies have lots of internal warnings :(
import warnings
warnings.filterwarnings("ignore", category=DeprecationWarning)
warnings.filterwarnings("ignore", category=RuntimeWarning) # TODO: remove this when google-crc32c publish a python3.12 wheel
import capnp
from aiohttp import web
if TYPE_CHECKING:
from aiortc.rtcdatachannel import RTCDataChannel
import aioice.ice
import signal
import threading
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from urllib.parse import urlparse, parse_qs
from typing import Any
from openpilot.system.webrtc.helpers import StreamRequestBody
from openpilot.system.webrtc.schema import generate_field
from openpilot.common.params import Params
from cereal import messaging, log
SESSION_TIMEOUT_SECONDS = 300
# socket trick: route lookup for 8.8.8.8 (nothing is sent or actually connected to)
# return the source interfaces IP which is the default interface of the device
@@ -41,20 +37,8 @@ def _default_route_ip() -> str | None:
finally:
s.close()
# aioice patch: gather ICE candidates only on the default-route interface
_get_host_addresses = aioice.ice.get_host_addresses
def _primary_host_addresses(use_ipv4: bool, use_ipv6: bool) -> list[str]:
addresses = _get_host_addresses(use_ipv4, use_ipv6)
primary = _default_route_ip()
if primary not in addresses:
return addresses
return [primary, ]
aioice.ice.get_host_addresses = _primary_host_addresses
class AsyncTaskRunner:
def __init__(self):
self.is_running = False
self.task = None
self.logger = logging.getLogger("webrtcd")
@@ -83,10 +67,10 @@ class CerealOutgoingMessageProxy(AsyncTaskRunner):
super().__init__()
self.services = list(services)
self.sm = messaging.SubMaster(self.services)
self.channels: list[RTCDataChannel] = []
self.channels = []
self._enabled = enabled
def add_channel(self, channel: 'RTCDataChannel'):
def add_channel(self, channel):
self.channels.append(channel)
def enable(self, enable: bool):
@@ -115,20 +99,17 @@ class CerealOutgoingMessageProxy(AsyncTaskRunner):
outgoing_msg = {"type": service, "logMonoTime": mono_time, "valid": valid, "data": msg_dict}
encoded_msg = json.dumps(outgoing_msg).encode()
for channel in self.channels:
if not channel.is_open():
continue
channel.send(encoded_msg)
async def run(self):
from aiortc.exceptions import InvalidStateError
while True:
if not self._enabled:
await asyncio.sleep(0.01)
continue
try:
self.update()
except InvalidStateError:
self.logger.warning("Cereal outgoing proxy invalid state (connection closed)")
break
except Exception:
self.logger.exception("Cereal outgoing proxy failure")
await asyncio.sleep(0.01)
@@ -169,17 +150,17 @@ class LivestreamBitrateController(AsyncTaskRunner):
high_level = 0.1 # drop immediately
med_level = 0.05 # drop after # of samples
low_level = 0 # raise after # of samples
down_samples = 5 # 1s
down_samples = 5
param_name = "LivestreamEncoderBitrate"
def __init__(self, peer_connection: Any, params: Params, enabled: bool = True):
def __init__(self, get_stats: Callable[[], dict[str, Any]], params: Params, enabled: bool = True):
super().__init__()
self.pc = peer_connection
self.get_stats = get_stats
self.params = params
self.level = 2
self._publish(self.bitrates[self.level])
self.prev_lost, self.prev_sent = None, None
self.prev_stats: tuple[Any, ...] | None = None
self.counter = 0
self.up_samples = 5 # 1s
self._auto = True
@@ -196,7 +177,7 @@ class LivestreamBitrateController(AsyncTaskRunner):
if not self._auto:
continue
loss_rate = await self._sample()
loss_rate = self._sample()
if loss_rate is None:
continue
if loss_rate >= self.med_level and self.level > 0:
@@ -213,22 +194,18 @@ class LivestreamBitrateController(AsyncTaskRunner):
self.counter = 0
self._publish(self.bitrates[self.level])
async def _sample(self) -> float | None:
report = await self.pc.getStats()
packets_lost = packets_sent = 0
for s in report.values():
if s.type == "remote-inbound-rtp":
packets_lost += s.packetsLost
elif s.type == "outbound-rtp":
packets_sent += s.packetsSent
if self.prev_lost is None:
self.prev_lost, self.prev_sent = packets_lost, packets_sent
def _sample(self) -> float | None:
report = next(iter(self.get_stats().values()), None)
if report is None:
return None
lost_delta = max(0, packets_lost - self.prev_lost)
sent_delta = max(0, packets_sent - self.prev_sent)
self.prev_lost, self.prev_sent = packets_lost, packets_sent
return lost_delta / sent_delta if sent_delta else 0.0
current = (report.ssrc, report.fraction_lost, report.packets_lost, report.highest_seq_no, report.jitter, report.lsr, report.dlsr)
if self.prev_stats == current:
return None
self.prev_stats = current
loss_rate = report.fraction_lost / 256
return loss_rate
def _publish(self, bitrate: float):
self.params.put(self.param_name, bitrate)
@@ -244,48 +221,41 @@ class LivestreamBitrateController(AsyncTaskRunner):
class StreamSession:
shared_pub_master = DynamicPubMaster([])
def __init__(self, body: StreamRequestBody, debug_mode: bool = False):
if debug_mode:
from aiortc.mediastreams import AudioStreamTrack, VideoStreamTrack
from aiortc.contrib.media import MediaBlackhole
def __init__(self, body: StreamRequestBody):
from openpilot.system.webrtc.device.video import LiveStreamVideoStreamTrack
from openpilot.system.webrtc.device.audio import AudioInputStreamTrack, AudioOutputSpeaker
from teleoprtc.builder import WebRTCAnswerBuilder
from teleoprtc.info import parse_info_from_offer
self.identifier = str(uuid.uuid4())
self.params = Params()
builder = WebRTCAnswerBuilder(body.sdp)
config = parse_info_from_offer(body.sdp)
builder = WebRTCAnswerBuilder(body.sdp, bind_address=_default_route_ip())
self.enabled = body.enabled
self.video_track = LiveStreamVideoStreamTrack(body.init_camera, self.enabled) if not debug_mode else VideoStreamTrack()
builder.add_video_stream(body.init_camera, self.video_track)
if config.expected_audio_track:
builder.add_audio_stream(AudioInputStreamTrack() if not debug_mode else AudioStreamTrack())
if config.incoming_audio_track:
self.audio_output_cls = AudioOutputSpeaker if not debug_mode else MediaBlackhole
builder.offer_to_receive_audio_stream()
self.video_tracks = []
for camera in body.cameras:
track = LiveStreamVideoStreamTrack(camera, self.enabled)
self.video_tracks.append(track)
builder.add_video_stream(camera, track)
self.stream = builder.stream()
self.is_body = "testJoystick" in body.bridge_services_in
self.incoming_bridge: CerealIncomingMessageProxy | None = None
self.incoming_bridge_services = body.bridge_services_in
self.outgoing_bridge: CerealOutgoingMessageProxy | None = None
self.bitrate_controller: LivestreamBitrateController | None = None
self.audio_output: AudioOutputSpeaker | MediaBlackhole | None = None
if len(body.bridge_services_in) > 0:
self.incoming_bridge = CerealIncomingMessageProxy(self.shared_pub_master)
if len(body.bridge_services_out) > 0:
self.outgoing_bridge = CerealOutgoingMessageProxy(body.bridge_services_out, self.enabled)
self.bitrate_controller = LivestreamBitrateController(self.stream.peer_connection, self.params, self.enabled)
self.bitrate_controller = LivestreamBitrateController(self.stream.get_receiver_report_stats, self.params, self.enabled)
self.run_task: asyncio.Task | None = None
self._cleanup_lock = asyncio.Lock()
self._cleanup_done = False
self.logger = logging.getLogger("webrtcd")
self.logger.info(
"New stream session (%s), init camera %s, video enabled %s, incoming services %s, outgoing services %s",
self.identifier, body.init_camera, body.enabled, body.bridge_services_in, body.bridge_services_out,
"New stream session (%s), video cameras %s, video enabled %s, incoming services %s, outgoing services %s",
self.identifier, [t.id for t in self.video_tracks], body.enabled, body.bridge_services_in, body.bridge_services_out,
)
def start(self):
@@ -310,16 +280,21 @@ class StreamSession:
match msg_type:
case "livestreamCameraSwitch":
self.video_track.switch_camera(payload["data"]["camera"])
# only needed for 1 track stream
if len(self.video_tracks) == 1:
self.video_tracks[0].switch_camera(payload["data"]["camera"])
case "livestreamSettings":
self.bitrate_controller.set_quality(payload["data"]["quality"])
if self.bitrate_controller is not None:
self.bitrate_controller.set_quality(payload["data"]["quality"])
case "livestreamVideoEnable":
enabled = payload["data"]["enabled"]
self.enabled = enabled
self.video_track.enable(enabled)
for track in self.video_tracks:
track.enable(enabled)
if self.outgoing_bridge is not None:
self.outgoing_bridge.enable(enabled)
self.bitrate_controller.enable(enabled)
if self.bitrate_controller is not None:
self.bitrate_controller.enable(enabled)
if not enabled:
self.params.put("LivestreamRequestKeyframe", True)
case "clockSync":
@@ -328,15 +303,29 @@ class StreamSession:
}})
self.stream.get_messaging_channel().send(pong)
case "enableTimingSei":
if hasattr(self.video_track, 'timing_sei_enabled'):
self.video_track.timing_sei_enabled = bool(payload["data"]["enabled"])
for track in self.video_tracks:
track.timing_sei_enabled = bool(payload["data"]["enabled"])
case _:
if payload.get("type") not in self.incoming_bridge_services:
if msg_type not in self.incoming_bridge_services:
return
self.incoming_bridge.send(message)
if self.incoming_bridge is not None:
self.incoming_bridge.send(message)
except Exception:
self.logger.exception("Cereal incoming proxy failure")
async def run_normal_session(self):
try:
await asyncio.wait_for(self.stream.wait_for_disconnection(), timeout=SESSION_TIMEOUT_SECONDS)
except TimeoutError:
self.logger.warning("Stream session (%s) timed out after %d s", self.identifier, SESSION_TIMEOUT_SECONDS)
try:
self.stream.get_messaging_channel().send(json.dumps({"type": "disconnect", "data": "Session timed out"}))
except Exception:
pass
async def run_body_session(self):
await self.stream.wait_for_disconnection()
async def run(self):
try:
self.params.put("LivestreamRequestKeyframe", True)
@@ -349,15 +338,14 @@ class StreamSession:
channel = self.stream.get_messaging_channel()
self.outgoing_bridge.add_channel(channel)
self.outgoing_bridge.start()
if self.stream.has_incoming_audio_track():
track = self.stream.get_incoming_audio_track(buffered=False)
self.audio_output = self.audio_output_cls()
self.audio_output.addTrack(track)
self.audio_output.start()
self.bitrate_controller.start()
if self.bitrate_controller is not None:
self.bitrate_controller.start()
self.logger.info("Stream session (%s) connected", self.identifier)
await self.stream.wait_for_disconnection()
if self.is_body:
await self.run_body_session()
else:
await self.run_normal_session()
self.logger.info("Stream session (%s) ended", self.identifier)
except Exception:
self.logger.exception("Stream session failure")
@@ -370,39 +358,52 @@ class StreamSession:
return
self._cleanup_done = True
self.params.put("LivestreamRequestKeyframe", False)
await self.bitrate_controller.stop()
if self.bitrate_controller is not None:
await self.bitrate_controller.stop()
if self.outgoing_bridge is not None:
await self.outgoing_bridge.stop()
if self.video_track is not None:
self.video_track.stop()
self.video_track = None
if self.audio_output is not None:
self.audio_output.stop()
self.audio_output = None
for track in self.video_tracks:
track.stop()
self.video_tracks.clear()
await self.stream.stop()
def schedule_teardown(app):
# if nothing connects for 5 seconds, tear down livestreaming processes
h = app.get('teardown')
if h:
h.cancel()
class ServerState:
def __init__(self):
self.streams: dict[str, StreamSession] = {}
self.stream_lock = asyncio.Lock()
self.teardown: asyncio.TimerHandle | None = None
# if nothing connects for 5 seconds, tear down livestreaming processes
def schedule_teardown(state: ServerState):
if state.teardown is not None:
state.teardown.cancel()
def clear():
if not app['streams']:
Params().put_bool("IsLiveStreaming", False)
app['teardown'] = asyncio.get_running_loop().call_later(5.0, clear)
if not state.streams:
Params().put_bool("IsLiveStreaming", False)
state.teardown = asyncio.get_running_loop().call_later(5.0, clear)
async def get_stream(request: 'web.Request'):
stream_dict, debug_mode = request.app['streams'], request.app['debug']
raw_body = await request.json()
body = StreamRequestBody(**raw_body)
def _json_response(obj: Any, status: int = 200) -> tuple[int, bytes, str]:
return (status, json.dumps(obj).encode(), "application/json; charset=utf-8")
async with request.app['stream_lock']:
def _text_response(text: str, status: int = 200) -> tuple[int, bytes, str]:
return (status, text.encode(), "text/plain; charset=utf-8")
async def handle_get_stream(state: ServerState, raw_body: bytes) -> tuple[int, bytes, str]:
stream_dict = state.streams
body = StreamRequestBody(**json.loads(raw_body))
async with state.stream_lock:
# don't remove existing connection on prewarm request
enabled = any(s.run_task and not s.run_task.done() and s.enabled for s in stream_dict.values())
if enabled and not body.enabled:
return web.json_response({"error": "busy", "message": "someone else is connected."})
return _json_response({"error": "busy", "message": "someone else is connected."})
for sid, s in list(stream_dict.items()):
if s.run_task and not s.run_task.done():
@@ -414,10 +415,15 @@ async def get_stream(request: 'web.Request'):
await s.stop()
stream_dict.pop(sid, None)
session = StreamSession(body, debug_mode)
session = StreamSession(body)
stream_dict[session.identifier] = session
try:
answer = await session.get_answer()
answer = await asyncio.wait_for(session.get_answer(), timeout=30)
except TimeoutError:
await session.stop()
stream_dict.pop(session.identifier, None)
logging.getLogger("webrtcd").exception("Timed out creating stream answer")
raise
except Exception:
await session.stop()
stream_dict.pop(session.identifier, None)
@@ -427,94 +433,189 @@ async def get_stream(request: 'web.Request'):
def remove_finished_session(_: asyncio.Task) -> None:
stream_dict.pop(session.identifier, None)
schedule_teardown(request.app)
schedule_teardown(state)
session.run_task.add_done_callback(remove_finished_session)
return web.json_response({"sdp": answer.sdp, "type": answer.type})
return _json_response({"sdp": answer.sdp, "type": answer.type})
async def get_schema(request: 'web.Request'):
services = request.query.get("services", "").split(",")
async def handle_get_schema(state: ServerState, services_param: str) -> tuple[int, bytes, str]:
services = services_param.split(",")
services = [s for s in services if s]
assert all(s in log.Event.schema.fields and not s.endswith("DEPRECATED") for s in services), "Invalid service name"
schema_dict = {s: generate_field(log.Event.schema.fields[s]) for s in services}
return web.json_response(schema_dict)
return _json_response(schema_dict)
async def post_notify(request: 'web.Request'):
try:
payload = await request.json()
except Exception as e:
raise web.HTTPBadRequest(text="Invalid JSON") from e
for session in list(request.app.get('streams', {}).values()):
async def handle_post_notify(state: ServerState, payload: Any) -> tuple[int, bytes, str]:
for session in list(state.streams.values()):
try:
ch = session.stream.get_messaging_channel()
ch.send(json.dumps(payload))
except Exception:
continue
return web.Response(status=200, text="OK")
return _text_response("OK")
async def on_shutdown(app: 'web.Application'):
for session in list(app['streams'].values()):
async def on_shutdown(state: ServerState):
for session in list(state.streams.values()):
try:
ch = session.stream.get_messaging_channel()
ch.send(json.dumps({"type": "disconnect", "data": "device streaming has been stopped."}))
except Exception:
pass
await session.stop()
del app['streams']
state.streams.clear()
@web.middleware
async def error_middleware(request: 'web.Request', handler):
try:
return await handler(request)
except Exception as e:
logging.getLogger("webrtcd").exception("Unhandled error handling %s", request.path)
return web.json_response({"error": "exception", "message": f"{type(e).__name__}: {e}"}, status=500)
class WebrtcdHandler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
# path -> allowed methods (aiohttp registered POST /stream, POST /notify, GET /schema + its auto HEAD)
_routes = {
"/schema": ("GET", "HEAD"),
"/stream": ("POST",),
"/notify": ("POST",),
}
def _send(self, status: int, body: bytes, content_type: str) -> None:
self.send_response(status)
self.send_header("Content-Type", content_type)
self.send_header("Content-Length", str(len(body)))
self.end_headers()
if self.command != "HEAD":
self.wfile.write(body)
def _read_body(self) -> bytes:
length = int(self.headers.get("Content-Length", 0))
return self.rfile.read(length) if length else b""
def _run(self, coro) -> tuple[int, bytes, str]:
return asyncio.run_coroutine_threadsafe(coro, self.server.loop).result()
def _dispatch_request(self) -> None:
parsed = urlparse(self.path)
allowed = self._routes.get(parsed.path)
try:
if allowed is None:
result = _json_response({"error": "not found"}, status=404)
elif self.command not in allowed:
result = _json_response({"error": "method not allowed"}, status=405)
elif parsed.path == "/schema":
services = parse_qs(parsed.query).get("services", [""])[0]
result = self._run(handle_get_schema(self.server.state, services))
elif parsed.path == "/stream":
result = self._run(handle_get_stream(self.server.state, self._read_body()))
else: # /notify
try:
payload = json.loads(self._read_body())
except Exception:
result = _json_response({"error": "bad request"}, status=400)
else:
result = self._run(handle_post_notify(self.server.state, payload))
except Exception as e:
logging.getLogger("webrtcd").exception("Unhandled error handling %s", self.path)
result = _json_response({"error": "exception", "message": f"{type(e).__name__}: {e}"}, status=500)
self._send(*result)
def do_GET(self) -> None:
self._dispatch_request()
def do_HEAD(self) -> None:
self._dispatch_request()
def do_POST(self) -> None:
self._dispatch_request()
def do_PUT(self) -> None:
self._dispatch_request()
def do_DELETE(self) -> None:
self._dispatch_request()
def do_PATCH(self) -> None:
self._dispatch_request()
def do_OPTIONS(self) -> None:
self._dispatch_request()
def log_message(self, format: str, *args: object) -> None: # noqa: A002 # stdlib override
# silence default access logging; errors are logged explicitly in _dispatch_request
pass
def prewarm_stream_session_imports(debug_mode: bool = False) -> None:
if debug_mode:
from aiortc.mediastreams import VideoStreamTrack
assert VideoStreamTrack
class WebrtcdHTTPServer(ThreadingHTTPServer):
daemon_threads = True
allow_reuse_address = True
state: ServerState
loop: asyncio.AbstractEventLoop
async def _shutdown(server: WebrtcdHTTPServer, state: ServerState, loop: asyncio.AbstractEventLoop) -> None:
# stop accepting new HTTP connections (blocks until serve_forever returns, so
# run it off the loop) then tear down active stream sessions.
await loop.run_in_executor(None, server.shutdown)
await on_shutdown(state)
loop.stop()
def prewarm_stream_session_imports() -> None:
from openpilot.system.webrtc.device.video import LiveStreamVideoStreamTrack
from teleoprtc.builder import WebRTCAnswerBuilder
assert LiveStreamVideoStreamTrack
assert WebRTCAnswerBuilder
def webrtcd_thread(host: str, port: int, debug: bool):
logging.basicConfig(level=logging.CRITICAL, handlers=[logging.StreamHandler()])
def webrtcd_thread(host: str, port: int):
logging.basicConfig(level=logging.INFO, handlers=[logging.StreamHandler()])
prewarm_start = time.monotonic()
prewarm_stream_session_imports(debug)
prewarm_stream_session_imports()
prewarm_end = time.monotonic()
logging.getLogger("webrtcd").info(f"webrtc prewarm finished in {(prewarm_end - prewarm_start) * 1000} ms")
app = web.Application(middlewares=[error_middleware])
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
state = ServerState()
app['streams'] = dict()
app['stream_lock'] = asyncio.Lock()
app['debug'] = debug
app.on_shutdown.append(on_shutdown)
app.router.add_post("/stream", get_stream)
app.router.add_post("/notify", post_notify)
app.router.add_get("/schema", get_schema)
server = WebrtcdHTTPServer((host, port), WebrtcdHandler)
server.state = state
server.loop = loop
web.run_app(app, host=host, port=port)
# serve HTTP on a daemon thread so the asyncio loop can own the main thread
http_thread = threading.Thread(target=server.serve_forever, name="webrtcd-http", daemon=True)
http_thread.start()
shutting_down = False
shutdown_task = None
def request_shutdown() -> None:
nonlocal shutting_down, shutdown_task
if shutting_down:
return
shutting_down = True
shutdown_task = loop.create_task(_shutdown(server, state, loop))
for sig in (signal.SIGINT, signal.SIGTERM):
loop.add_signal_handler(sig, request_shutdown)
try:
loop.run_forever()
finally:
server.server_close()
loop.close()
def main():
parser = argparse.ArgumentParser(description="WebRTC daemon")
parser.add_argument("--host", type=str, default="0.0.0.0", help="Host to listen on")
parser.add_argument("--port", type=int, default=5001, help="Port to listen on")
parser.add_argument("--debug", action="store_true", help="Enable debug mode")
args = parser.parse_args()
webrtcd_thread(args.host, args.port, args.debug)
webrtcd_thread(args.host, args.port)
if __name__=="__main__":
+15 -14
View File
@@ -1,9 +1,7 @@
import abc
from typing import Dict, List
from typing import Dict, List, Optional
import aiortc
from teleoprtc.stream import WebRTCBaseStream, WebRTCOfferStream, WebRTCAnswerStream, ConnectionProvider
from teleoprtc.stream import RTCSessionDescription, WebRTCBaseStream, WebRTCOfferStream, WebRTCAnswerStream, ConnectionProvider
from teleoprtc.tracks import TiciVideoStreamTrack, TiciTrackWrapper
@@ -14,11 +12,12 @@ class WebRTCStreamBuilder(abc.ABC):
class WebRTCOfferBuilder(WebRTCStreamBuilder):
def __init__(self, connection_provider: ConnectionProvider):
def __init__(self, connection_provider: ConnectionProvider, bind_address: Optional[str] = None):
self.connection_provider = connection_provider
self.bind_address = bind_address
self.requested_camera_types: List[str] = []
self.requested_audio = False
self.audio_tracks: List[aiortc.MediaStreamTrack] = []
self.audio_tracks: List[object] = []
self.messaging_enabled = False
def offer_to_receive_video_stream(self, camera_type: str):
@@ -28,7 +27,7 @@ class WebRTCOfferBuilder(WebRTCStreamBuilder):
def offer_to_receive_audio_stream(self):
self.requested_audio = True
def add_audio_stream(self, track: aiortc.MediaStreamTrack):
def add_audio_stream(self, track: object):
assert len(self.audio_tracks) == 0
self.audio_tracks = [track]
@@ -43,32 +42,34 @@ class WebRTCOfferBuilder(WebRTCStreamBuilder):
video_producer_tracks=[],
audio_producer_tracks=self.audio_tracks,
should_add_data_channel=self.messaging_enabled,
bind_address=self.bind_address,
)
class WebRTCAnswerBuilder(WebRTCStreamBuilder):
def __init__(self, offer_sdp: str):
def __init__(self, offer_sdp: str, bind_address: Optional[str] = None):
self.offer_sdp = offer_sdp
self.video_tracks: Dict[str, aiortc.MediaStreamTrack] = dict()
self.bind_address = bind_address
self.video_tracks: Dict[str, TiciVideoStreamTrack] = {}
self.requested_audio = False
self.audio_tracks: List[aiortc.MediaStreamTrack] = []
self.audio_tracks: List[object] = []
def offer_to_receive_audio_stream(self):
self.requested_audio = True
def add_video_stream(self, camera_type: str, track: aiortc.MediaStreamTrack):
def add_video_stream(self, camera_type: str, track: object):
assert camera_type not in self.video_tracks
assert camera_type in ["driver", "wideRoad", "road"]
if not isinstance(track, TiciVideoStreamTrack):
track = TiciTrackWrapper(camera_type, track)
self.video_tracks[camera_type] = track
def add_audio_stream(self, track: aiortc.MediaStreamTrack):
def add_audio_stream(self, track: object):
assert len(self.audio_tracks) == 0
self.audio_tracks = [track]
def stream(self) -> WebRTCBaseStream:
description = aiortc.RTCSessionDescription(sdp=self.offer_sdp, type="offer")
description = RTCSessionDescription(sdp=self.offer_sdp, type="offer")
return WebRTCAnswerStream(
description,
consumed_camera_types=[],
@@ -76,5 +77,5 @@ class WebRTCAnswerBuilder(WebRTCStreamBuilder):
video_producer_tracks=list(self.video_tracks.values()),
audio_producer_tracks=self.audio_tracks,
should_add_data_channel=False,
bind_address=self.bind_address,
)
+50
View File
@@ -0,0 +1,50 @@
import dataclasses
import struct
from typing import List
@dataclasses.dataclass(frozen=True)
class RtcpReceiverReport:
ssrc: int
fraction_lost: int
packets_lost: int
highest_seq_no: int
jitter: int
lsr: int
dlsr: int
def _decode_receiver_reports(message: bytes) -> List[RtcpReceiverReport]:
reports: List[RtcpReceiverReport] = []
offset = 0
while offset + 4 <= len(message):
flags, packet_type, length_words = struct.unpack_from("!BBH", message, offset)
packet_end = offset + (length_words + 1) * 4
if flags >> 6 != 2 or packet_end > len(message):
break
report_count = flags & 0x1F
if packet_type == 200: # Sender Report
report_offset = offset + 28
elif packet_type == 201: # Receiver Report
report_offset = offset + 8
else:
offset = packet_end
continue
if report_offset + report_count * 24 > packet_end:
break
for i in range(report_count):
block_offset = report_offset + i * 24
ssrc, loss, highest_seq_no, jitter, lsr, dlsr = struct.unpack_from("!IIIIII", message, block_offset)
fraction_lost = loss >> 24
packets_lost = loss & 0xFFFFFF
if packets_lost & 0x800000:
packets_lost -= 1 << 24
reports.append(RtcpReceiverReport(ssrc, fraction_lost, packets_lost, highest_seq_no, jitter, lsr, dlsr))
offset = packet_end
return reports
+19 -9
View File
@@ -1,6 +1,6 @@
import dataclasses
import aiortc
from libdatachannel import Description
@dataclasses.dataclass
@@ -15,13 +15,23 @@ def parse_info_from_offer(sdp: str) -> StreamingMediaInfo:
"""
helper function to parse info about outgoing and incoming streams from an offer sdp
"""
desc = aiortc.sdp.SessionDescription.parse(sdp)
audio_tracks = [m for m in desc.media if m.kind == "audio"]
video_tracks = [m for m in desc.media if m.kind == "video" and m.direction in ["recvonly", "sendrecv"]]
application_tracks = [m for m in desc.media if m.kind == "application"]
has_incoming_audio_track = next((t for t in audio_tracks if t.direction in ["sendonly", "sendrecv"]), None) is not None
has_incoming_datachannel = len(application_tracks) > 0
expects_outgoing_audio_track = next((t for t in audio_tracks if t.direction in ["recvonly", "sendrecv"]), None) is not None
desc = Description(sdp, Description.Type.Offer)
n_video = 0
expected_audio_track = False
incoming_audio_track = False
incoming_datachannel = desc.has_application()
return StreamingMediaInfo(len(video_tracks), expects_outgoing_audio_track, has_incoming_audio_track, has_incoming_datachannel)
for i in range(desc.media_count()):
media = desc.media(i)
if media is None:
continue
direction = media.direction()
if media.type() == "video" and direction in (Description.Direction.RecvOnly, Description.Direction.SendRecv):
n_video += 1
elif media.type() == "audio":
if direction in (Description.Direction.RecvOnly, Description.Direction.SendRecv):
expected_audio_track = True
if direction in (Description.Direction.SendOnly, Description.Direction.SendRecv):
incoming_audio_track = True
return StreamingMediaInfo(n_video, expected_audio_track, incoming_audio_track, incoming_datachannel)
+299 -156
View File
@@ -1,13 +1,29 @@
import abc
import asyncio
import contextlib
import dataclasses
import logging
from typing import Any, Awaitable, Callable, Dict, List, Optional
import random
from typing import Any, Awaitable, Callable, Dict, List, Optional, Tuple, Union
import aiortc
from aiortc.contrib.media import MediaRelay
from libdatachannel import (
Configuration,
DataChannel,
Description,
FrameInfo,
H264RtpPacketizer,
IceServer,
NalUnit,
PeerConnection,
PliHandler,
RtcpNackResponder,
RtcpSrReporter,
RtpPacketizationConfig,
Track,
)
from teleoprtc.tracks import parse_video_track_id
from teleoprtc.decoder import RtcpReceiverReport, _decode_receiver_reports
from teleoprtc.tracks import TiciVideoStreamTrack, parse_video_track_id
@dataclasses.dataclass
@@ -16,138 +32,238 @@ class StreamingOffer:
video: List[str]
ConnectionProvider = Callable[[StreamingOffer], Awaitable[aiortc.RTCSessionDescription]]
MessageHandler = Callable[[bytes], Awaitable[None]]
@dataclasses.dataclass
class RTCSessionDescription:
sdp: str
type: str
ConnectionProvider = Callable[[StreamingOffer], Awaitable[RTCSessionDescription]]
MessageHandler = Callable[[Union[bytes, str]], None]
class WebRTCBaseStream(abc.ABC):
# destorying wrapper on close can cause deadlock
# TODO: upstream a fix to this
_retained_messaging_channels: List[DataChannel] = []
_retain_messaging_channel_on_close = False
def __init__(self,
consumed_camera_types: List[str],
consume_audio: bool,
video_producer_tracks: List[aiortc.MediaStreamTrack],
audio_producer_tracks: List[aiortc.MediaStreamTrack],
should_add_data_channel: bool):
self.peer_connection = aiortc.RTCPeerConnection()
self.media_relay = MediaRelay()
video_producer_tracks: List[TiciVideoStreamTrack],
audio_producer_tracks: List[Any],
should_add_data_channel: bool,
bind_address: Optional[str] = None):
config = Configuration()
config.force_media_transport = True
config.disable_auto_negotiation = True
config.ice_servers = [IceServer("stun:stun.l.google.com:19302")]
if bind_address is not None:
config.bind_address = bind_address
self.peer_connection = PeerConnection(config)
self.expected_incoming_camera_types = consumed_camera_types
self.expected_incoming_audio = consume_audio
self.expected_number_of_incoming_media: Optional[int] = None
self.incoming_camera_tracks: Dict[str, aiortc.MediaStreamTrack] = dict()
self.incoming_audio_tracks: List[aiortc.MediaStreamTrack] = []
self.outgoing_video_tracks: List[aiortc.MediaStreamTrack] = video_producer_tracks
self.outgoing_audio_tracks: List[aiortc.MediaStreamTrack] = audio_producer_tracks
self.incoming_camera_tracks: Dict[str, Any] = {}
self.incoming_audio_tracks: List[Any] = []
self.outgoing_video_tracks = video_producer_tracks
self.outgoing_audio_tracks = audio_producer_tracks
self.should_add_data_channel = should_add_data_channel
self.messaging_channel: Optional[aiortc.RTCDataChannel] = None
self.messaging_channel: Optional[DataChannel] = None
self.incoming_message_handlers: List[MessageHandler] = []
self._consumer_tracks: List[Track] = []
self._sender_tasks: List[asyncio.Task] = []
self._track_state: List[Tuple[Track, TiciVideoStreamTrack, RtpPacketizationConfig]] = []
self._receiver_reports: Dict[str, RtcpReceiverReport] = {}
self._receiver_report_tracks: Dict[str, Tuple[Track, int]] = {}
self.incoming_media_ready_event = asyncio.Event()
self.messaging_channel_ready_event = asyncio.Event()
self.connection_attempted_event = asyncio.Event()
self.connection_stopped_event = asyncio.Event()
self.gathering_complete_event = asyncio.Event()
self._loop: Optional[asyncio.AbstractEventLoop] = None
self.peer_connection.on("connectionstatechange", self._on_connectionstatechange)
self.peer_connection.on("datachannel", self._on_incoming_datachannel)
self.peer_connection.on("track", self._on_incoming_track)
self.peer_connection.on_state_change(self._on_connectionstatechange)
self.peer_connection.on_gathering_state_change(self._on_gatheringstatechange)
self.peer_connection.on_data_channel(self._on_incoming_datachannel)
if self.expected_incoming_camera_types or self.expected_incoming_audio:
self.peer_connection.on_track(self._on_incoming_track)
self.logger = logging.getLogger("WebRTCStream")
def _log_debug(self, msg: Any, *args):
self.logger.debug(f"{type(self)}() {msg}", *args)
def _call_soon_threadsafe(self, fn: Callable, *args) -> None:
if self._loop is not None and self._loop.is_running():
self._loop.call_soon_threadsafe(fn, *args)
else:
fn(*args)
def _set_event(self, event: asyncio.Event) -> None:
self._call_soon_threadsafe(event.set)
@property
def _number_of_incoming_media(self) -> int:
media = len(self.incoming_camera_tracks) + len(self.incoming_audio_tracks)
# if stream does not add data_channel, then it means its incoming
media += int(self.messaging_channel is not None) if not self.should_add_data_channel else 0
return media
def _add_consumer_transceivers(self):
for _ in self.expected_incoming_camera_types:
self.peer_connection.addTransceiver("video", direction="recvonly")
for camera_type in self.expected_incoming_camera_types:
media = Description.Video(camera_type, Description.Direction.RecvOnly)
media.add_h264_codec(96)
track = self.peer_connection.add_track(media)
self._consumer_tracks.append(track)
self.incoming_camera_tracks[camera_type] = track
if self.expected_incoming_audio:
self.peer_connection.addTransceiver("audio", direction="recvonly")
media = Description.Audio("audio", Description.Direction.RecvOnly)
media.add_opus_codec(111)
track = self.peer_connection.add_track(media)
self._consumer_tracks.append(track)
self.incoming_audio_tracks.append(track)
def _find_trackless_transceiver(self, kind: str) -> Optional[aiortc.RTCRtpTransceiver]:
transceivers = self.peer_connection.getTransceivers()
target_transceiver = None
for t in transceivers:
if t.kind == kind and t.sender.track is None:
target_transceiver = t
break
def _find_offer_video(self, remote_sdp: str, used_mids: set[str]) -> Tuple[str, int]:
desc = Description(remote_sdp, Description.Type.Offer)
for i in range(desc.media_count()):
media = desc.media(i)
if media is None or media.type() != "video" or media.mid() in used_mids:
continue
for payload_type in media.payload_types():
with contextlib.suppress(ValueError):
rtp_map = media.rtp_map(payload_type)
if rtp_map is not None and rtp_map.format.upper() == "H264":
return media.mid(), payload_type
raise ValueError("Remote SDP does not offer H264 video")
return target_transceiver
def _make_video_media(self, track: TiciVideoStreamTrack, remote_sdp: str, used_mids: set[str]) -> Tuple[Description.Video, int, int, str]:
mid, payload_type = self._find_offer_video(remote_sdp, used_mids)
used_mids.add(mid)
ssrc = random.randint(1, 0xFFFFFFFF)
cname = f"teleoprtc-{random.getrandbits(32):08x}"
stream_id = f"stream-{random.getrandbits(32):08x}"
media = Description.Video(mid, Description.Direction.SendOnly)
media.add_h264_codec(payload_type)
media.add_ssrc(ssrc, cname, stream_id, track.id)
return media, ssrc, payload_type, cname
def _add_producer_tracks(self):
def _add_producer_tracks(self, remote_sdp: Optional[str] = None):
used_mids: set[str] = set()
for track in self.outgoing_video_tracks:
target_transceiver = self._find_trackless_transceiver(track.kind)
if target_transceiver is None:
self.peer_connection.addTransceiver(track.kind, direction="sendonly")
media, ssrc, payload_type, cname = self._make_video_media(track, remote_sdp or "", used_mids)
rtc_track = self.peer_connection.add_track(media)
sender = self.peer_connection.addTrack(track)
if hasattr(track, "codec_preference") and track.codec_preference() is not None:
transceiver = next(t for t in self.peer_connection.getTransceivers() if t.sender == sender)
self._force_codec(transceiver, track.codec_preference(), "video")
for track in self.outgoing_audio_tracks:
target_transceiver = self._find_trackless_transceiver(track.kind)
if target_transceiver is None:
self.peer_connection.addTransceiver(track.kind, direction="sendonly")
rtp_config = RtpPacketizationConfig(ssrc, cname, payload_type, H264RtpPacketizer.CLOCK_RATE)
rtp_config.start_timestamp = random.randint(0, 0xFFFFFFFF)
rtp_config.timestamp = rtp_config.start_timestamp
rtp_config.sequence_number = random.randint(0, 0xFFFF)
self.peer_connection.addTrack(track)
packetizer = H264RtpPacketizer(NalUnit.Separator.LongStartSequence, rtp_config, 1200)
packetizer.add_to_chain(RtcpSrReporter(rtp_config))
packetizer.add_to_chain(PliHandler(track.request_keyframe))
packetizer.add_to_chain(RtcpNackResponder())
rtc_track.set_media_handler(packetizer)
def _add_messaging_channel(self, channel: Optional[aiortc.RTCDataChannel] = None):
if not channel:
channel = self.peer_connection.createDataChannel("data", ordered=True)
camera_type, _ = parse_video_track_id(track.id)
rtc_track.reset_callbacks()
self._receiver_report_tracks[camera_type] = (rtc_track, ssrc)
self._track_state.append((rtc_track, track, rtp_config))
for handler in self.incoming_message_handlers:
channel.on("message", handler)
if self.outgoing_audio_tracks:
raise NotImplementedError("Audio producer tracks are not implemented with libdatachannel")
if channel.readyState == "open":
self.messaging_channel_ready_event.set()
else:
channel.on("open", lambda: self.messaging_channel_ready_event.set())
def _add_messaging_channel(self, channel: Optional[DataChannel] = None):
if channel is None:
channel = self.peer_connection.create_data_channel("data")
self.messaging_channel = channel
def _force_codec(self, transceiver: aiortc.RTCRtpTransceiver, codec: str, stream_type: str):
codec_mime = f"{stream_type}/{codec.upper()}"
rtp_codecs = aiortc.RTCRtpSender.getCapabilities(stream_type).codecs
rtp_codec = [c for c in rtp_codecs if c.mimeType == codec_mime]
transceiver.setCodecPreferences(rtp_codec)
def on_message(message: Union[bytes, str]):
for handler in list(self.incoming_message_handlers):
self._call_soon_threadsafe(handler, message)
def _on_connectionstatechange(self):
self._log_debug("connection state is %s", self.peer_connection.connectionState)
if self.peer_connection.connectionState in ['connected', 'failed']:
self.connection_attempted_event.set()
if self.peer_connection.connectionState in ['disconnected', 'closed', 'failed']:
self.connection_stopped_event.set()
def on_open():
self._set_event(self.messaging_channel_ready_event)
def _on_incoming_track(self, track: aiortc.MediaStreamTrack):
self._log_debug("got track: %s %s", track.kind, track.id)
if track.kind == "video":
camera_type, _ = parse_video_track_id(track.id)
if camera_type in self.expected_incoming_camera_types:
self.incoming_camera_tracks[camera_type] = track
elif track.kind == "audio":
if self.expected_incoming_audio:
self.incoming_audio_tracks.append(track)
def on_closed():
self._set_event(self.connection_stopped_event)
channel.on_message(on_message)
channel.on_open(on_open)
channel.on_closed(on_closed)
if channel.is_open():
self._set_event(self.messaging_channel_ready_event)
self._on_after_media()
def _on_incoming_datachannel(self, channel: aiortc.RTCDataChannel):
self._log_debug("got data channel: %s", channel.label)
if channel.label == "data" and self.messaging_channel is None:
def _retain_messaging_channel(self) -> None:
if self.messaging_channel is None:
return
# No native callback can be running before a remote description is set.
if not self.messaging_channel_ready_event.is_set() and self.peer_connection.remote_description() is None:
self.messaging_channel = None
return
if self._retain_messaging_channel_on_close:
self._retained_messaging_channels.append(self.messaging_channel)
self.messaging_channel = None
def _on_connectionstatechange(self, state: PeerConnection.State):
self._log_debug("connection state is %s", state)
if state in (PeerConnection.State.Connected, PeerConnection.State.Failed):
self._set_event(self.connection_attempted_event)
if state in (PeerConnection.State.Disconnected, PeerConnection.State.Closed, PeerConnection.State.Failed):
self._set_event(self.connection_stopped_event)
def _on_gatheringstatechange(self, state: PeerConnection.GatheringState):
self._log_debug("gathering state is %s", state)
if state == PeerConnection.GatheringState.Complete:
self._set_event(self.gathering_complete_event)
def _on_incoming_track(self, track: Track):
self._log_debug("got track: %s", track.mid())
try:
camera_type, _ = parse_video_track_id(track.mid())
except ValueError:
camera_type = track.mid()
if camera_type in self.expected_incoming_camera_types:
self.incoming_camera_tracks[camera_type] = track
elif self.expected_incoming_audio:
self.incoming_audio_tracks.append(track)
self._on_after_media()
def _on_incoming_datachannel(self, channel: DataChannel):
self._log_debug("got data channel: %s", channel.label())
if channel.label() == "data" and self.messaging_channel is None:
self._add_messaging_channel(channel)
self._on_after_media()
def _update_receiver_report(self, camera_type: str, ssrc: int, message: bytes) -> None:
for report in _decode_receiver_reports(message):
if report.ssrc == ssrc:
self._receiver_reports[camera_type] = report
def _on_after_media(self):
if self._number_of_incoming_media == self.expected_number_of_incoming_media:
self.incoming_media_ready_event.set()
if self.expected_number_of_incoming_media is not None and self._number_of_incoming_media >= self.expected_number_of_incoming_media:
self._set_event(self.incoming_media_ready_event)
def _parse_incoming_streams(self, remote_sdp: str):
desc = aiortc.sdp.SessionDescription.parse(remote_sdp)
audio_video_media_count = len([m for m in desc.media if m.kind in ["audio", "video"] and m.direction in ["sendonly", "sendrecv"]])
data_media_count = int(any(m for m in desc.media if m.kind == "application")) if not self.should_add_data_channel else 0
self.expected_number_of_incoming_media = audio_video_media_count + data_media_count
desc = Description(remote_sdp, Description.Type.Offer)
media_count = 0
for i in range(desc.media_count()):
media = desc.media(i)
if media is None:
continue
direction = media.direction()
if media.type() in ("audio", "video") and direction in (Description.Direction.SendOnly, Description.Direction.SendRecv):
media_count += 1
data_media_count = int(desc.has_application()) if not self.should_add_data_channel else 0
self.expected_number_of_incoming_media = media_count + data_media_count
if self.expected_number_of_incoming_media == 0:
self._set_event(self.incoming_media_ready_event)
def has_incoming_video_track(self, camera_type: str) -> bool:
return camera_type in self.incoming_camera_tracks
@@ -158,65 +274,122 @@ class WebRTCBaseStream(abc.ABC):
def has_messaging_channel(self) -> bool:
return self.messaging_channel is not None
def get_incoming_video_track(self, camera_type: str, buffered: bool = False) -> aiortc.MediaStreamTrack:
def get_incoming_video_track(self, camera_type: str) -> Track:
assert camera_type in self.incoming_camera_tracks, "Video tracks are not enabled on this stream"
assert self.is_started, "Stream must be started"
return self.incoming_camera_tracks[camera_type]
track = self.incoming_camera_tracks[camera_type]
relay_track = self.media_relay.subscribe(track, buffered=buffered)
return relay_track
def get_incoming_audio_track(self, buffered: bool = False) -> aiortc.MediaStreamTrack:
def get_incoming_audio_track(self) -> Track:
assert len(self.incoming_audio_tracks) > 0, "Audio tracks are not enabled on this stream"
assert self.is_started, "Stream must be started"
return self.incoming_audio_tracks[0]
track = self.incoming_audio_tracks[0]
relay_track = self.media_relay.subscribe(track, buffered=buffered)
return relay_track
def get_messaging_channel(self) -> aiortc.RTCDataChannel:
def get_messaging_channel(self) -> DataChannel:
assert self.messaging_channel is not None, "Messaging channel is not enabled on this stream"
assert self.is_started, "Stream must be started"
return self.messaging_channel
def get_receiver_report_stats(self) -> Dict[str, RtcpReceiverReport]:
return dict(self._receiver_reports)
def set_message_handler(self, message_handler: MessageHandler):
self.incoming_message_handlers.append(message_handler)
if self.messaging_channel is not None:
self.messaging_channel.on("message", message_handler)
@property
def is_started(self) -> bool:
return self.peer_connection is not None and \
self.peer_connection.localDescription is not None and \
self.peer_connection.remoteDescription is not None and \
self.peer_connection.connectionState != "closed"
self.peer_connection.local_description() is not None and \
self.peer_connection.remote_description() is not None and \
self.peer_connection.state() != PeerConnection.State.Closed
@property
def is_connected_and_ready(self) -> bool:
return self.peer_connection is not None and \
self.peer_connection.connectionState == "connected" and \
self.peer_connection.state() == PeerConnection.State.Connected and \
(self.expected_number_of_incoming_media == 0 or self.incoming_media_ready_event.is_set())
async def _wait_for_gathering_complete(self):
if self.peer_connection.gathering_state() != PeerConnection.GatheringState.Complete:
await self.gathering_complete_event.wait()
async def _send_track_loop(self, rtc_track: Track, producer_track: TiciVideoStreamTrack, rtp_config: RtpPacketizationConfig):
while True:
if not rtc_track.is_open():
await asyncio.sleep(0.01)
continue
try:
packet = await producer_track.recv()
data = bytes(packet)
if not data:
continue
pts = int(packet.pts or 0)
timestamp = (rtp_config.start_timestamp + pts) & 0xFFFFFFFF
rtc_track.send_frame(data, FrameInfo(timestamp))
except asyncio.CancelledError:
raise
except Exception:
self.logger.exception("Error in send track loop for track %s", producer_track.id)
self._set_event(self.connection_stopped_event)
break
async def _receiver_report_loop(self):
while True:
for camera_type, (rtc_track, ssrc) in self._receiver_report_tracks.items():
for _ in range(32):
try:
message = rtc_track.receive()
if message is None: # go until queue empty (bounded to 32)
break
if isinstance(message, bytes):
self._update_receiver_report(camera_type, ssrc, message)
except asyncio.CancelledError:
raise
except Exception:
self.logger.exception("Error receiving report for %s", camera_type)
break
await asyncio.sleep(0.05)
def _start_sender_tasks(self):
for rtc_track, producer_track, rtp_config in self._track_state:
self._sender_tasks.append(asyncio.create_task(self._send_track_loop(rtc_track, producer_track, rtp_config)))
if self._track_state:
self._sender_tasks.append(asyncio.create_task(self._receiver_report_loop()))
async def wait_for_connection(self):
assert self.is_started
await self.connection_attempted_event.wait()
if self.peer_connection.connectionState != 'connected':
if self.peer_connection.state() != PeerConnection.State.Connected:
raise ValueError("Connection failed.")
if self.expected_number_of_incoming_media:
await self.incoming_media_ready_event.wait()
if self.messaging_channel is not None:
await self.messaging_channel_ready_event.wait()
self._start_sender_tasks()
async def wait_for_disconnection(self):
assert self.is_connected_and_ready, "Stream is not connected/ready yet (make sure wait_for_connection was awaited)"
await self.connection_stopped_event.wait()
async def stop(self):
await self.peer_connection.close()
for task in self._sender_tasks:
task.cancel()
for task in self._sender_tasks:
with contextlib.suppress(asyncio.CancelledError):
await task
self._sender_tasks.clear()
self._retain_messaging_channel()
self.peer_connection.close()
self.incoming_camera_tracks.clear()
self.incoming_audio_tracks.clear()
self._consumer_tracks.clear()
self._track_state.clear()
self._receiver_reports.clear()
self._receiver_report_tracks.clear()
@abc.abstractmethod
async def start(self) -> aiortc.RTCSessionDescription:
async def start(self) -> RTCSessionDescription:
raise NotImplementedError
@@ -225,76 +398,46 @@ class WebRTCOfferStream(WebRTCBaseStream):
super().__init__(*args, **kwargs)
self.session_provider = session_provider
async def start(self) -> aiortc.RTCSessionDescription:
async def start(self) -> RTCSessionDescription:
self._loop = asyncio.get_running_loop()
self._add_consumer_transceivers()
if self.should_add_data_channel:
self._add_messaging_channel()
self._add_producer_tracks()
offer = await self.peer_connection.createOffer()
await self.peer_connection.setLocalDescription(offer)
actual_offer = self.peer_connection.localDescription
self.peer_connection.set_local_description(Description.Type.Offer)
await self._wait_for_gathering_complete()
actual_offer = self.peer_connection.local_description()
streaming_offer = StreamingOffer(
sdp=actual_offer.sdp,
sdp=str(actual_offer),
video=list(self.expected_incoming_camera_types),
)
remote_answer = await self.session_provider(streaming_offer)
self._parse_incoming_streams(remote_sdp=remote_answer.sdp)
await self.peer_connection.setRemoteDescription(remote_answer)
actual_answer = self.peer_connection.remoteDescription
self.peer_connection.set_remote_description(Description(remote_answer.sdp, Description.Type.Answer))
self._on_after_media()
actual_answer = self.peer_connection.remote_description()
return actual_answer
return RTCSessionDescription(str(actual_answer), actual_answer.type_string())
class WebRTCAnswerStream(WebRTCBaseStream):
def __init__(self, session: aiortc.RTCSessionDescription, *args, **kwargs):
_retain_messaging_channel_on_close = True
def __init__(self, session: RTCSessionDescription, *args, **kwargs):
super().__init__(*args, **kwargs)
self.session = session
def _probe_video_codecs(self) -> List[str]:
codecs = []
for track in self.outgoing_video_tracks:
if hasattr(track, "codec_preference") and track.codec_preference() is not None:
codecs.append(track.codec_preference())
return codecs
def _override_incoming_video_codecs(self, remote_sdp: str, codecs: List[str]) -> str:
desc = aiortc.sdp.SessionDescription.parse(remote_sdp)
codec_mimes = [f"video/{c}" for c in codecs]
for m in desc.media:
if m.kind != "video":
continue
preferred_codecs: List[aiortc.RTCRtpCodecParameters] = [c for c in m.rtp.codecs if c.mimeType in codec_mimes]
if len(preferred_codecs) == 0:
raise ValueError(f"None of {preferred_codecs} codecs is supported in remote SDP")
m.rtp.codecs = preferred_codecs
m.fmt = [c.payloadType for c in preferred_codecs]
return str(desc)
async def start(self) -> aiortc.RTCSessionDescription:
assert self.peer_connection.remoteDescription is None, "Connection already established"
self._add_consumer_transceivers()
# since we sent already encoded frames in some cases (e.g. livestream video tracks are in H264), we need to force aiortc to actually use it
# we do that by overriding supported codec information on incoming sdp
preferred_codecs = self._probe_video_codecs()
if len(preferred_codecs) > 0:
self.session.sdp = self._override_incoming_video_codecs(self.session.sdp, preferred_codecs)
async def start(self) -> RTCSessionDescription:
self._loop = asyncio.get_running_loop()
assert self.peer_connection.remote_description() is None, "Connection already established"
self._parse_incoming_streams(remote_sdp=self.session.sdp)
await self.peer_connection.setRemoteDescription(self.session)
self.peer_connection.set_remote_description(Description(self.session.sdp, Description.Type.Offer))
self._add_producer_tracks(self.session.sdp)
self._add_producer_tracks()
answer = await self.peer_connection.createAnswer()
await self.peer_connection.setLocalDescription(answer)
actual_answer = self.peer_connection.localDescription
return actual_answer
self.peer_connection.set_local_description(Description.Type.Answer)
await self._wait_for_gathering_complete()
actual_answer = self.peer_connection.local_description()
return RTCSessionDescription(str(actual_answer), actual_answer.type_string())
+29 -35
View File
@@ -1,11 +1,11 @@
import asyncio
import logging
import time
import fractions
from typing import Any, Optional, Tuple
import logging
import uuid
from typing import Any, Tuple
import aiortc
from aiortc.mediastreams import VIDEO_CLOCK_RATE, VIDEO_TIME_BASE
VIDEO_CLOCK_RATE = 90000
VIDEO_TIME_BASE = fractions.Fraction(1, VIDEO_CLOCK_RATE)
def video_track_id(camera_type: str, track_id: str) -> str:
@@ -21,57 +21,51 @@ def parse_video_track_id(track_id: str) -> Tuple[str, str]:
return camera_type, track_id
class TiciVideoStreamTrack(aiortc.MediaStreamTrack):
class TiciVideoStreamTrack:
"""
Abstract video track which associates video track with camera_type
Abstract video track which associates video track with camera_type.
"""
kind = "video"
def __init__(self, camera_type: str, dt: float, time_base: fractions.Fraction = VIDEO_TIME_BASE, clock_rate: int = VIDEO_CLOCK_RATE):
assert camera_type in ["driver", "wideRoad", "road"]
super().__init__()
# override track id to include camera type - client needs that for identification
self._id: str = video_track_id(camera_type, self._id)
self._dt: float = dt
self._id: str = video_track_id(camera_type, str(uuid.uuid4()))
self._time_base: fractions.Fraction = time_base
self._clock_rate: int = clock_rate
self._start: Optional[float] = None
self._logger = logging.getLogger("WebRTCStream")
self.readyState = "live"
@property
def id(self) -> str:
return self._id
def stop(self) -> None:
self.readyState = "ended"
def log_debug(self, msg: Any, *args):
self._logger.debug(f"{type(self)}() {msg}", *args)
async def next_pts(self, current_pts) -> float:
pts: float = current_pts + self._dt * self._clock_rate
async def recv(self):
raise NotImplementedError()
data_time = pts * self._time_base
if self._start is None:
self._start = time.time() - data_time
else:
wait_time = self._start + data_time - time.time()
await asyncio.sleep(wait_time)
return pts
def codec_preference(self) -> Optional[str]:
return None
def request_keyframe(self) -> None:
pass
class TiciTrackWrapper(aiortc.MediaStreamTrack):
class TiciTrackWrapper(TiciVideoStreamTrack):
"""
Associates video track with camera_type
Associates a generic video track with camera_type.
"""
def __init__(self, camera_type: str, track: aiortc.MediaStreamTrack):
def __init__(self, camera_type: str, track: Any):
assert track.kind == "video"
assert not isinstance(track, TiciVideoStreamTrack)
super().__init__()
super().__init__(camera_type, getattr(track, "_dt", 0.05))
self._id = video_track_id(camera_type, track.id)
self._track = track
@property
def kind(self) -> str:
return self._track.kind
async def recv(self):
return await self._track.recv()
def stop(self) -> None:
super().stop()
if hasattr(self._track, "stop"):
self._track.stop()
+76 -42
View File
@@ -2,12 +2,45 @@
set -euo pipefail
HOST="${1:-comma@192.168.3.110}"
IMAGE="${2:-/Users/dominickthompson/Desktop/system8.img.xz}"
IMAGE="${2:-/Users/dominickthompson/Desktop/system14.img.xz}"
METADATA="${3:-${IMAGE}.metadata.json}"
SCRIPT_DIR="$(cd -- "$(dirname -- "${BASH_SOURCE[0]}")" && pwd)"
REPO_ROOT="$(cd -- "${SCRIPT_DIR}/../.." && pwd)"
SSH_KEY="${SSH_KEY:-${REPO_ROOT}/system/hardware/tici/id_rsa}"
SSH_OPTS=(-i "$SSH_KEY" -o BatchMode=yes -o StrictHostKeyChecking=no -o UserKnownHostsFile=/dev/null)
SSH_OPTS=(-o BatchMode=yes -o StrictHostKeyChecking=no -o UserKnownHostsFile=/dev/null)
if [[ -n "${SSH_KEY:-}" ]]; then
SSH_OPTS=(-i "$SSH_KEY" -o IdentitiesOnly=yes "${SSH_OPTS[@]}")
fi
for required_path in "$IMAGE" "$METADATA"; do
[[ -f "$required_path" ]] || { echo "missing file: $required_path" >&2; exit 1; }
done
metadata_value() {
python3 - "$METADATA" "$1" <<'PY'
import json
import sys
from pathlib import Path
metadata = json.loads(Path(sys.argv[1]).read_text(encoding="utf-8"))
value = metadata[sys.argv[2]]
if isinstance(value, (dict, list)):
raise SystemExit(f"metadata field {sys.argv[2]} is not scalar")
print(value)
PY
}
BASE_VERSION="$(metadata_value base_version)"
EXPECTED_VERSION="$(metadata_value target_version)"
RAW_HASH="$(metadata_value raw_sha256)"
RAW_SIZE="$(metadata_value raw_size)"
EXPECTED_XZ_HASH="$(metadata_value xz_sha256)"
ACTUAL_XZ_HASH="$(shasum -a 256 "$IMAGE" | awk '{print $1}')"
[[ "$ACTUAL_XZ_HASH" == "$EXPECTED_XZ_HASH" ]] || {
echo "compressed image hash mismatch: got $ACTUAL_XZ_HASH, expected $EXPECTED_XZ_HASH" >&2
exit 1
}
SESSION="local_agnos_flash"
REMOTE_DIR="/data/local_agnos_flash"
@@ -15,40 +48,49 @@ REMOTE_MANIFEST="${REMOTE_DIR}/agnos-local-system.json"
REMOTE_RUNNER="${REMOTE_DIR}/run_flash.sh"
REMOTE_AGNOS="${REMOTE_DIR}/agnos.py"
PORT="8989"
EXPECTED_VERSION="12.8.28"
RAW_HASH="4c01245932068aedfceb41cb1aab1f7f044f6659aa2fe2de558f99e2d3aa5793"
RAW_SIZE="5368709120"
if [[ ! -f "$IMAGE" ]]; then
echo "missing image: $IMAGE" >&2
exit 1
fi
IMAGE_NAME="$(basename "$IMAGE")"
REMOTE_IMAGE="${REMOTE_DIR}/${IMAGE_NAME}"
INSTALLED_VERSION="$(ssh "${SSH_OPTS[@]}" "$HOST" 'tr -d "\r\n" </VERSION')"
case "$INSTALLED_VERSION" in
19.6|19.6.*) ;;
*)
echo "refusing flash: device is on incompatible AGNOS $INSTALLED_VERSION; candidate is based on upstream $BASE_VERSION" >&2
exit 1
;;
esac
echo "[CHECK] Device AGNOS: $INSTALLED_VERSION"
echo "[CHECK] Candidate AGNOS: $EXPECTED_VERSION"
echo "[CHECK] Candidate XZ hash: $ACTUAL_XZ_HASH"
ssh "${SSH_OPTS[@]}" "$HOST" "mkdir -p '$REMOTE_DIR'"
scp "${SSH_OPTS[@]}" "$IMAGE" "$HOST:$REMOTE_IMAGE"
LOCAL_AGNOS="$(mktemp "${TMPDIR:-/tmp}/agnos-local.XXXXXX.py")"
trap 'rm -f "$LOCAL_AGNOS"' EXIT
python3 - "$REPO_ROOT/system/hardware/tici/agnos.py" "$LOCAL_AGNOS" <<'PY'
import sys
from pathlib import Path
src, dst = map(Path, sys.argv[1:])
data = src.read_text(encoding="utf-8")
needle = "import openpilot.system.updated.casync.casync as casync"
if data.count(needle) != 1:
raise SystemExit("could not isolate the unused casync dependency in agnos.py")
data = data.replace(
"import openpilot.system.updated.casync.casync as casync",
"""class _UnusedCasync:
needle,
'''class _UnusedCasync:
ChunkReader = object
ChunkDict = object
def __getattr__(self, name):
raise RuntimeError("casync support is unavailable in local AGNOS flash runner")
casync = _UnusedCasync()""",
casync = _UnusedCasync()''',
)
dst.write_text(data, encoding="utf-8")
PY
scp "${SSH_OPTS[@]}" "$LOCAL_AGNOS" "$HOST:$REMOTE_AGNOS"
rm -f "$LOCAL_AGNOS"
ssh "${SSH_OPTS[@]}" "$HOST" "cat > '$REMOTE_MANIFEST'" <<MANIFEST
[
@@ -73,40 +115,33 @@ set -euo pipefail
: "${REMOTE_MANIFEST:?}"
: "${REMOTE_AGNOS:?}"
: "${PORT:?}"
: "${IMAGE_NAME:?}"
: "${EXPECTED_VERSION:?}"
exec > >(tee -a "${REMOTE_DIR}/flash.log") 2>&1
echo "[STEP] Local AGNOS system flash"
echo "[CHECK] Installed AGNOS: $(cat /VERSION 2>/dev/null || echo unknown)"
echo "[CHECK] Target AGNOS: ${EXPECTED_VERSION}"
echo "[CHECK] Active slot: $(abctl --boot_slot)"
df -h /data
if [[ ! -f "$REMOTE_AGNOS" ]]; then
echo "[ERROR] $REMOTE_AGNOS not found" >&2
exit 1
fi
if [[ -x /usr/local/venv/bin/python3 ]]; then
PYTHON_BIN="/usr/local/venv/bin/python3"
else
PYTHON_BIN="python3"
fi
PYTHON_BIN="/usr/local/venv/bin/python3"
[[ -x "$PYTHON_BIN" ]] || { echo "[ERROR] managed Python is unavailable" >&2; exit 1; }
pkill -f "http.server ${PORT}.*${REMOTE_DIR}" >/dev/null 2>&1 || true
"$PYTHON_BIN" -m http.server "$PORT" --bind 127.0.0.1 --directory "$REMOTE_DIR" >"${REMOTE_DIR}/http.log" 2>&1 &
http_pid="$!"
trap 'kill "$http_pid" >/dev/null 2>&1 || true' EXIT
http_ready=0
for _ in $(seq 1 20); do
if "$PYTHON_BIN" - "${PORT}" "${IMAGE_NAME}" <<'PY'
if "$PYTHON_BIN" - "$REMOTE_MANIFEST" <<'PY'
import json
import sys
import urllib.request
from pathlib import Path
port, image_name = sys.argv[1], sys.argv[2]
with urllib.request.urlopen(f"http://127.0.0.1:{port}/{image_name}", timeout=2) as resp:
resp.read(1)
for entry in json.loads(Path(sys.argv[1]).read_text(encoding="utf-8")):
with urllib.request.urlopen(entry["url"], timeout=2) as response:
response.read(1)
PY
then
http_ready=1
@@ -115,24 +150,23 @@ PY
sleep 0.25
done
if [[ "$http_ready" != "1" ]]; then
echo "[ERROR] Local image HTTP server did not become ready" >&2
[[ "${http_ready:-0}" == "1" ]] || {
echo "[ERROR] local image server did not become ready" >&2
cat "${REMOTE_DIR}/http.log" >&2 || true
exit 1
fi
}
echo "[FLASH] Flashing local system image to inactive AGNOS slot"
echo "[FLASH] Writing and verifying the candidate in the inactive system slot"
PYTHONPATH="$(dirname "$REMOTE_AGNOS")" "$PYTHON_BIN" "$REMOTE_AGNOS" --swap "$REMOTE_MANIFEST"
echo "[DONE] AGNOS flashed and slot swapped"
echo "[REBOOT] Rebooting now"
echo "[DONE] Candidate written, verified, and selected"
sudo reboot
REMOTE_RUNNER
ssh "${SSH_OPTS[@]}" "$HOST" "tmux kill-session -t '$SESSION' >/dev/null 2>&1 || true"
ssh "${SSH_OPTS[@]}" "$HOST" "rm -f '$REMOTE_DIR/flash.log' '$REMOTE_DIR/http.log'"
ssh "${SSH_OPTS[@]}" "$HOST" \
"tmux new-session -d -s '$SESSION' \"REMOTE_DIR='$REMOTE_DIR' REMOTE_MANIFEST='$REMOTE_MANIFEST' REMOTE_AGNOS='$REMOTE_AGNOS' PORT='$PORT' IMAGE_NAME='$IMAGE_NAME' EXPECTED_VERSION='$EXPECTED_VERSION' bash '$REMOTE_RUNNER'\""
"tmux new-session -d -s '$SESSION' \"REMOTE_DIR='$REMOTE_DIR' REMOTE_MANIFEST='$REMOTE_MANIFEST' REMOTE_AGNOS='$REMOTE_AGNOS' PORT='$PORT' EXPECTED_VERSION='$EXPECTED_VERSION' bash '$REMOTE_RUNNER'\""
echo "Started remote tmux session: $SESSION"
echo "Watch it with: ssh $HOST 'tmux attach -t $SESSION'"
echo "After reboot, run tools/agnos/validate_agnos_runtime.sh $EXPECTED_VERSION on the device."
+564 -1295
View File
File diff suppressed because it is too large Load Diff
+75 -86
View File
@@ -1,106 +1,95 @@
import importlib.util
import json
from pathlib import Path
import runpy
import pytest
from tools.agnos.patch_system_reset_image import (
AMDGPU_FIRMWARE_SHA256,
COMMA_SH_DISPLAY_WAIT_PATCH_MARKER,
comma_sh_has_expected_display_wait,
find_default_reference_manifest,
format_debugfs_mode,
patch_comma_sh_display_wait,
patch_setup_branding_script,
sha256_zstd_payload,
)
def _load_patch_module():
path = Path(__file__).resolve().parent / "patch_system_reset_image.py"
spec = importlib.util.spec_from_file_location("patch_system_reset_image_under_test", path)
assert spec is not None and spec.loader is not None
module = importlib.util.module_from_spec(spec)
spec.loader.exec_module(module)
return module
ORIGINAL_DISPLAY_WAIT = b'''#!/usr/bin/env bash
echo "waiting for magic"
for i in {1..200}; do
if systemctl is-active --quiet magic && [ -S /tmp/drmfd.sock ]; then
break
fi
sleep 0.1
done
if systemctl is-active --quiet magic && [ -S /tmp/drmfd.sock ]; then
echo "magic ready after ${SECONDS}s"
else
echo "timed out waiting for magic, ${SECONDS}s"
fi
exec /data/continue.sh
'''
patch_image = _load_patch_module()
def test_patch_comma_sh_display_wait_uses_available_display_service():
patched = patch_comma_sh_display_wait(ORIGINAL_DISPLAY_WAIT)
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_only_version_and_additive_runtime_packages_are_mutable():
assert patch_image.ALLOWED_IMAGE_MUTATIONS == {
patch_image.VERSION_PATH_IN_IMAGE,
*patch_image.STAR_PILOT_DEPENDENCY_PATHS,
}
assert set(patch_image.C3_DEPENDENCY_PATHS) < set(patch_image.STAR_PILOT_DEPENDENCY_PATHS)
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_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/installer",
"/usr/comma/setup",
"/usr/comma/comma.sh",
"/usr/comma/magic.py",
}
@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
@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(("mode", "expected"), [
("100775", "0100775"),
("100644", "0100644"),
("040755", "040755"),
("120777", "0120777"),
])
def test_format_debugfs_mode(mode, expected):
assert format_debugfs_mode(mode) == expected
@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_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_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_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_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)
def test_default_reference_manifest_uses_sibling_openpilot(tmp_path):
primary = tmp_path / "starpilot/system/hardware/tici/agnos.json"
reference = tmp_path / "openpilot/openpilot/system/hardware/tici/agnos.json"
primary.parent.mkdir(parents=True)
reference.parent.mkdir(parents=True)
primary.write_text("[]")
reference.write_text("[]")
assert find_default_reference_manifest(primary) == reference.resolve()
def test_protected_payload_validation_reports_any_drift():
patch_image.validate_protected_payloads(dict(patch_image.PROTECTED_PAYLOAD_HASHES))
changed = dict(patch_image.PROTECTED_PAYLOAD_HASHES)
changed["/usr/comma/reset"] = "0" * 64
with pytest.raises(RuntimeError, match="/usr/comma/reset"):
patch_image.validate_protected_payloads(changed)
+55
View File
@@ -0,0 +1,55 @@
#!/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' </VERSION)"
if [[ -n "$EXPECTED_VERSION" && "$actual_version" != "$EXPECTED_VERSION" ]]; then
fail "AGNOS version is $actual_version, expected $EXPECTED_VERSION"
fi
site_count="$(find "$SITE_PACKAGES" -mindepth 1 -maxdepth 1 -print | wc -l | tr -d ' ')"
[[ "$site_count" == "253" ]] || fail "managed venv has $site_count site-packages entries, expected upstream 213 plus 40 additive StarPilot dependencies"
echo "[CHECK] AGNOS version: $actual_version"
echo "[CHECK] managed venv entries: $site_count"
"$PYTHON_BIN" -c 'import aiohttp, capnp, crcmod, Crypto, cv2, jsonrpc, kaitaistruct, mapbox_earcut, numpy, onnx, pyaudio, raylib, serial, tqdm, xattr; print("[CHECK] runtime imports: ok")'
if [[ -d "$REPO_ROOT" ]]; then
(
cd "$REPO_ROOT"
"$PYTHON_BIN" -c 'from openpilot.system.manager import manager; print("[CHECK] manager import: ok")'
)
else
echo "[CHECK] manager import: deferred until software is installed"
fi
sha256sum --check --strict <<'HASHES'
779db62d2d4c5f8ce504c5d1f2994d34a9f35296d5efb7f3a48cb1e8a0d4778e /etc/NetworkManager/NetworkManager.conf
45e653e2f709c027fad41f2d86b70e008b72c6bf4d34590b4765ebe8fe3ea948 /etc/NetworkManager/conf.d/10-globally-managed-devices.conf
fb33a80bf8c78b3af004d4b294c47a0139e37742c1d0d5a6a7663c7d1f4a2b48 /lib/systemd/system/NetworkManager.service
9df4edbeb5849de03f9c2d691d04646af84a3ef74c2f33be8e73d9281daebe99 /usr/comma/updater
97ed6413515d0674442c42ae6e20baccf66dd6bb4ec382ee4cf0cc5ebe84e739 /usr/comma/reset
c4416e66b127b31c17d08e6723ad46d12af7683e56626512d66c708b5f347ac9 /usr/comma/magic.py
85f6d9e54286a3842920d6967b187478b4e43d6171c331d72d3fb3102106e101 /usr/comma/installer
934f74ab4b2ac06048418c2857be3a041e192ec03c09979987691c23c91353bd /usr/comma/setup_keys
c382ce266653bad781c25e403ddea4af508aa6f3ea2eef3f568d964982fad9d6 /usr/comma/setup
bcba2b336cf0ca852786f8a58bbce407e0e9fe952c26fc5d903f6d9a34b44b4f /usr/comma/comma.sh
HASHES
[[ "$(systemctl is-enabled NetworkManager)" == "enabled" ]] || fail "NetworkManager is not enabled"
[[ "$(systemctl is-active NetworkManager)" == "active" ]] || fail "NetworkManager is not active"
echo "[PASS] AGNOS runtime, manager, recovery payloads, and networking validated"
@@ -16,6 +16,7 @@ from pathlib import Path
REQUIRED_DIRS = [
("usr/local/lib", "/usr/local/lib"),
("usr/local/include", "/usr/local/include"),
("usr/local/venv/lib/python3.12/site-packages/raylib/install", "/usr/local/venv/lib/python3.12/site-packages/raylib/install"),
("lib/aarch64-linux-gnu", "/lib/aarch64-linux-gnu"),
("usr/lib/aarch64-linux-gnu", "/usr/lib/aarch64-linux-gnu"),
("usr/include", "/usr/include"),
Generated
+39
View File
@@ -335,6 +335,26 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/d1/d6/3965ed04c63042e047cb6a3e6ed1a63a35087b6a609aa3a15ed8ac56c221/colorama-0.4.6-py2.py3-none-any.whl", hash = "sha256:4f1d9991f5acc0ca119f9d443620b77f9d6b33703e51011c16baf57afb285fc6", size = 25335, upload-time = "2022-10-25T02:36:20.889Z" },
]
[[package]]
name = "comma-deps-capnproto"
version = "1.0.1.post98"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/ca/83/d3e6346a31491be1d378e4585f37a7979eb772018616abfa74fb27750f1e/comma_deps_capnproto-1.0.1.post98-py3-none-macosx_11_0_arm64.whl", hash = "sha256:9f4d08682df92411b360bec855cb6475990313cd0ecd8ed5c6ee02befb9db913", size = 2407343, upload-time = "2026-07-23T17:01:21.247Z" },
{ url = "https://files.pythonhosted.org/packages/b0/8b/6f2a29d50ed4c8741dbf0a34ab109899268d09753518cd693e881bbf1a9d/comma_deps_capnproto-1.0.1.post98-py3-none-manylinux_2_28_aarch64.whl", hash = "sha256:6cdf838a8d415ac71e1f52306624ab3ab6f27777f6ce89c0059c4129ea7b7f62", size = 2506355, upload-time = "2026-07-23T17:01:25.254Z" },
{ url = "https://files.pythonhosted.org/packages/08/24/e91f2203d62e4db9de7dae06dd0cdefb1e000d8b3ba0bde48367be7e5b63/comma_deps_capnproto-1.0.1.post98-py3-none-manylinux_2_28_x86_64.whl", hash = "sha256:3d95993c9aff0c89e39ca965e995021dd3dccdfc3d5d85152916cf4bf651b7ec", size = 2590764, upload-time = "2026-07-23T17:01:29.062Z" },
]
[[package]]
name = "comma-deps-ffmpeg"
version = "7.1.0.post98"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/20/59/4899ac0fa54905f43e237fff122008d6c591918b41e36880eeb18cd6279c/comma_deps_ffmpeg-7.1.0.post98-py3-none-macosx_11_0_arm64.whl", hash = "sha256:7816c5adc9c6a7462209ccf1d023c42d7e38a4eee47aa2f43d45dfc8320063f8", size = 7326312, upload-time = "2026-07-23T17:01:55.975Z" },
{ url = "https://files.pythonhosted.org/packages/29/cb/6e047c19c39977c5ae322ad698b91d8d9fce43314cb86563de91bb161982/comma_deps_ffmpeg-7.1.0.post98-py3-none-manylinux_2_28_aarch64.whl", hash = "sha256:45ad401b4058e3f7efb8d6841e1f187e8c265e942e56e7fb01db1ba9096e6b78", size = 4437675, upload-time = "2026-07-23T17:01:59.971Z" },
{ url = "https://files.pythonhosted.org/packages/76/3d/cda4b19fa5a7b26921a518143c94fd3632030a34b157cab6d6f10f2c86bc/comma_deps_ffmpeg-7.1.0.post98-py3-none-manylinux_2_28_x86_64.whl", hash = "sha256:9e7034739d45a45254c4200a555d343b22ce117429bebc2de2c1b05f050bfc8c", size = 4681499, upload-time = "2026-07-23T17:02:03.963Z" },
]
[[package]]
name = "contourpy"
version = "1.3.3"
@@ -905,6 +925,19 @@ wheels = [
{ url = "https://files.pythonhosted.org/packages/da/e9/0d4add7873a73e462aeb45c036a2dead2562b825aa46ba326727b3f31016/kiwisolver-1.4.9-pp311-pypy311_pp73-win_amd64.whl", hash = "sha256:fb940820c63a9590d31d88b815e7a3aa5915cad3ce735ab45f0c730b39547de1", size = 73929, upload-time = "2025-08-10T21:27:48.236Z" },
]
[[package]]
name = "libdatachannel-py"
version = "2026.1.0.dev2"
source = { registry = "https://pypi.org/simple" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/46/2f/68e8306327ddef4b2133d2efb163cb05b319759ce8bd50b8b32dcd03dd95/libdatachannel_py-2026.1.0.dev2-cp312-cp312-macosx_15_0_arm64.whl", hash = "sha256:6607fa1439e1b5bfceecd387c433470c9d45e439c3c06fa064f5c4669ad7e582", size = 1213155, upload-time = "2026-05-19T03:37:12.796Z" },
{ url = "https://files.pythonhosted.org/packages/fc/e3/10aed36ffaf1744795322aae612db777991575b72a9f04e2c677c2c022bf/libdatachannel_py-2026.1.0.dev2-cp312-cp312-macosx_26_0_arm64.whl", hash = "sha256:a060b1250f57d1fccb36e3a6b36ac8f4fd34926a6b51c564e787e8b7206458aa", size = 1224706, upload-time = "2026-05-19T03:37:12.679Z" },
{ url = "https://files.pythonhosted.org/packages/09/a9/103fc647a8f9c721ab140fe8d2f8dbd90817e917ca763c4eb0f843fe247e/libdatachannel_py-2026.1.0.dev2-cp312-cp312-manylinux_2_35_aarch64.whl", hash = "sha256:0b8aa3be2fa3654ea24f882756d6599e847de866a44f7a291c60994548a2debb", size = 1638879, upload-time = "2026-05-19T03:37:18.138Z" },
{ url = "https://files.pythonhosted.org/packages/09/a8/0dc7d3fe80fc247ec165dbd455bc2a1a307ef65702a43e24473202c2bf42/libdatachannel_py-2026.1.0.dev2-cp312-cp312-manylinux_2_35_x86_64.whl", hash = "sha256:339a79fcbc8c6caf91c620f6e4a0f1b8ccb6a941a966d10a3135c980ae4651a6", size = 1718006, upload-time = "2026-05-19T03:37:11.891Z" },
{ url = "https://files.pythonhosted.org/packages/6d/86/30904a8753e9db60d8c3cf8efda09585fc68f2004d3d7aa2910c93a8eed5/libdatachannel_py-2026.1.0.dev2-cp312-cp312-manylinux_2_38_aarch64.whl", hash = "sha256:b9f476cb065b50856ab2e53bf774ccca9c6a66454b7ce903aaf6dfcdef2a4482", size = 1643757, upload-time = "2026-05-19T03:37:12.207Z" },
{ url = "https://files.pythonhosted.org/packages/d7/9d/1e10131396d28e84a8088a63c14978cc215f6677dc85acdd96b6068f0664/libdatachannel_py-2026.1.0.dev2-cp312-cp312-manylinux_2_38_x86_64.whl", hash = "sha256:1f31db7347549edcd69fcc1ecb8b31e7183894808ccf9387afc49a4d68debaae", size = 1751748, upload-time = "2026-05-19T03:37:06.98Z" },
]
[[package]]
name = "libusb1"
version = "3.3.1"
@@ -1582,6 +1615,8 @@ dependencies = [
{ name = "aiortc" },
{ name = "casadi" },
{ name = "cffi" },
{ name = "comma-deps-capnproto", marker = "python_full_version >= '3.12'" },
{ name = "comma-deps-ffmpeg", marker = "python_full_version >= '3.12'" },
{ name = "crcmod" },
{ name = "cython" },
{ name = "future-fstrings" },
@@ -1589,6 +1624,7 @@ dependencies = [
{ name = "jeepney" },
{ name = "json-rpc" },
{ name = "kaitaistruct" },
{ name = "libdatachannel-py", marker = "python_full_version >= '3.12'" },
{ name = "libusb1" },
{ name = "mapbox-earcut" },
{ name = "numpy" },
@@ -1681,6 +1717,8 @@ requires-dist = [
{ name = "casadi", specifier = ">=3.6.6" },
{ name = "cffi" },
{ name = "codespell", marker = "extra == 'testing'" },
{ name = "comma-deps-capnproto", marker = "python_full_version >= '3.12'" },
{ name = "comma-deps-ffmpeg", marker = "python_full_version >= '3.12'" },
{ name = "coverage", marker = "extra == 'testing'" },
{ name = "crcmod" },
{ name = "cython" },
@@ -1694,6 +1732,7 @@ requires-dist = [
{ name = "jinja2", marker = "extra == 'docs'" },
{ name = "json-rpc" },
{ name = "kaitaistruct" },
{ name = "libdatachannel-py", marker = "python_full_version >= '3.12'", specifier = ">=2026.1.0.dev2" },
{ name = "libusb1" },
{ name = "mapbox-earcut" },
{ name = "matplotlib", marker = "extra == 'dev'" },