test/hil: avoid parallel MTP probe races

This commit is contained in:
Zixun LI
2026-07-28 23:27:54 +02:00
parent d8595dafcd
commit b868d6d268

View File

@ -39,6 +39,7 @@
import argparse import argparse
import io import io
import itertools import itertools
import math
import os import os
import random import random
import re import re
@ -64,7 +65,7 @@ _mp = multiprocessing.get_context('fork')
Pool, Lock, Semaphore, Manager = _mp.Pool, _mp.Lock, _mp.Semaphore, _mp.Manager Pool, Lock, Semaphore, Manager = _mp.Pool, _mp.Lock, _mp.Semaphore, _mp.Manager
import hashlib import hashlib
import ctypes import ctypes
from pymtp import MTP from pymtp import LIBMTP_DeviceEntry, LIBMTP_RawDevice, MTP
import string import string
# --- per-board dev-session locks (see test/hil/board_lock.py) ------------ # --- per-board dev-session locks (see test/hil/board_lock.py) ------------
@ -127,11 +128,11 @@ def enum_timeout() -> int:
return _enum_timeout return _enum_timeout
def wait_until(predicate, step: float = 1.0): def wait_until(predicate, step: float = 1.0, timeout: float | None = None):
"""Poll predicate under the per-attempt enum budget. Deadline-based so a slow predicate """Poll predicate under the per-attempt enum budget. Deadline-based so a slow predicate
body (subprocess, libmtp scan) counts against the budget. Returns the first truthy body (subprocess, libmtp scan) counts against the budget. An explicit timeout overrides
predicate value, or None on timeout.""" that budget. Returns the first truthy predicate value, or None on timeout."""
deadline = time.monotonic() + enum_timeout() deadline = time.monotonic() + (enum_timeout() if timeout is None else timeout)
while True: while True:
r = predicate() r = predicate()
if r: if r:
@ -503,23 +504,120 @@ def read_disk_file(uid: str, lun: int, fname: str) -> bytes:
return data return data
def open_mtp_dev(uid): def open_mtp_dev(uid: str):
mtp = MTP() mtp = MTP()
last_usb = None
deadline = time.monotonic() + 2 * enum_timeout()
def try_open(): def find_usb():
# unmount gio/gvfs MTP mount which blocks libmtp from accessing the device nonlocal last_usb
subprocess.run(f"gio mount -u mtp://TinyUsb_TinyUsb_Device_{uid}/", for serial_fname in glob.glob('/sys/bus/usb/devices/*/serial'):
shell=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL) dev_path = Path(serial_fname).parent
for raw in mtp.detect_devices(): try:
mtp.device = mtp.mtp.LIBMTP_Open_Raw_Device(ctypes.byref(raw)) if (Path(serial_fname).read_text().strip().lower() != uid.lower()
if mtp.device: or (dev_path / 'idVendor').read_text().strip() != 'cafe'
sn = mtp.get_serialnumber().decode('utf-8') or (dev_path / 'idProduct').read_text().strip() != '4017'):
if sn == uid: continue
return mtp busnum = int((dev_path / 'busnum').read_text())
mtp.disconnect() devnum = int((dev_path / 'devnum').read_text())
last_usb = (dev_path.name, busnum, devnum)
usb_node = Path('/dev/bus/usb') / f'{busnum:03d}' / f'{devnum:03d}'
if usb_node.exists():
return dev_path, busnum, devnum
except (OSError, ValueError):
pass
return None return None
return wait_until(try_open) def remaining() -> float:
return max(0.0, deadline - time.monotonic())
target = wait_until(find_usb, step=0.05, timeout=remaining())
if target is None:
if last_usb:
name, busnum, devnum = last_usb
raise AssertionError(
f'MTP USB node not ready for {uid} at {name} ({busnum:03d}/{devnum:03d})')
raise AssertionError(f'MTP USB device not enumerated for {uid}')
dev_path, busnum, devnum = target
wait_seconds = max(1, math.ceil(remaining()))
try:
udev_wait = subprocess.run(
['udevadm', 'wait', '--initialized=yes', f'--timeout={wait_seconds}', str(dev_path)],
capture_output=True, text=True, timeout=wait_seconds + 2)
except FileNotFoundError:
udev_wait = None
except subprocess.TimeoutExpired as e:
raise AssertionError(
f'udev initialization timed out for MTP {uid} at {busnum:03d}/{devnum:03d}') from e
if udev_wait is not None and udev_wait.returncode != 0:
try:
wait_help = subprocess.run(
['udevadm', 'wait', '--help'], stdout=subprocess.DEVNULL,
stderr=subprocess.DEVNULL, timeout=2)
wait_supported = wait_help.returncode == 0
except (FileNotFoundError, subprocess.TimeoutExpired):
wait_supported = False
if wait_supported:
detail = (udev_wait.stderr or udev_wait.stdout).strip().replace('\n', ' ')
detail = detail[-300:] or 'no diagnostic'
raise AssertionError(
f'udevadm wait failed for MTP {uid} at {busnum:03d}/{devnum:03d}: {detail}')
udev_wait = None
if udev_wait is None:
# systemd < 251 has no target-specific udev wait. Its libmtp rule creates
# this link only after synchronous mtp-probe has released the interface.
def find_libmtp_marker():
found = find_usb()
if found is None:
return None
found_path, found_busnum, found_devnum = found
marker = Path('/dev') / f'libmtp-{found_path.name}'
usb_node = Path('/dev/bus/usb') / f'{found_busnum:03d}' / f'{found_devnum:03d}'
if marker.exists() and marker.resolve() == usb_node:
return found
return None
target = wait_until(find_libmtp_marker, step=0.05, timeout=remaining())
if target is None:
raise AssertionError(
f'udevadm wait unsupported and libmtp marker absent for MTP {uid}; '
'install libmtp-runtime')
dev_path, busnum, devnum = target
elif find_usb() != target:
raise AssertionError(f'MTP USB device {uid} changed while waiting for udev initialization')
# A desktop GVFS session may claim MTP after udev probing. This is a no-op on
# headless runners, but preserves support for rigs where the mount exists.
try:
subprocess.run(['gio', 'mount', '-u', f'mtp://TinyUsb_TinyUsb_Device_{uid}/'],
stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, timeout=2)
except (FileNotFoundError, subprocess.TimeoutExpired):
pass
# TinyUSB needs no libmtp device quirks. Construct its raw entry directly so
# this test never probes another MTP board that is still being initialized.
entry = LIBMTP_DeviceEntry(None, 0xcafe, None, 0x4017, 0)
raw = LIBMTP_RawDevice(entry, busnum, devnum)
mtp.device = mtp.mtp.LIBMTP_Open_Raw_Device(ctypes.byref(raw))
if not mtp.device:
raise AssertionError(f'libmtp could not open MTP {uid} at {busnum:03d}/{devnum:03d}')
try:
serial_raw = mtp.get_serialnumber()
serial = serial_raw.decode('utf-8') if serial_raw else ''
if serial.lower() != uid.lower():
raise AssertionError(f'MTP serial mismatch at {busnum:03d}/{devnum:03d}: {serial}')
except Exception:
try:
mtp.disconnect()
except Exception:
pass
raise
return mtp
def get_printer_dev(id: str, vendor_str, product_str, ifnum: int): def get_printer_dev(id: str, vendor_str, product_str, ifnum: int):
@ -1432,15 +1530,13 @@ def test_device_mtp(board):
_null = os.open(os.devnull, os.O_WRONLY) _null = os.open(os.devnull, os.O_WRONLY)
os.dup2(_null, fd) os.dup2(_null, fd)
mtp = open_mtp_dev(uid) try:
mtp = open_mtp_dev(uid)
# --- AFTER: restore stderr --- finally:
os.dup2(_saved, fd) # --- AFTER: restore stderr ---
os.close(_null) os.dup2(_saved, fd)
os.close(_saved) os.close(_null)
os.close(_saved)
if mtp is None or mtp.device is None:
assert False, 'MTP device not found'
try: try:
assert b"TinyUSB" == mtp.get_manufacturer(), 'MTP wrong manufacturer' assert b"TinyUSB" == mtp.get_manufacturer(), 'MTP wrong manufacturer'