mirror of
https://github.com/firestar5683/StarPilot.git
synced 2026-08-25 02:03:43 +08:00
Compare commits
1 Commits
Dom
..
Dom_New_Agnos
| Author | SHA1 | Date | |
|---|---|---|---|
| b56eca9428 |
+32
-2
@@ -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
@@ -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
|
||||
|
||||
@@ -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
@@ -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()
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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"
|
||||
}
|
||||
]
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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""
|
||||
|
||||
@@ -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
@@ -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__":
|
||||
|
||||
@@ -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,
|
||||
)
|
||||
|
||||
|
||||
@@ -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
|
||||
@@ -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
@@ -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())
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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."
|
||||
|
||||
Regular → Executable
+564
-1295
File diff suppressed because it is too large
Load Diff
@@ -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)
|
||||
|
||||
Executable
+55
@@ -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"),
|
||||
|
||||
@@ -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'" },
|
||||
|
||||
Reference in New Issue
Block a user