diff --git a/test/hil/hil_test.py b/test/hil/hil_test.py index 6ba756a77..7bc3e0868 100755 --- a/test/hil/hil_test.py +++ b/test/hil/hil_test.py @@ -39,6 +39,7 @@ import argparse import io import itertools +import math import os import random import re @@ -64,7 +65,7 @@ _mp = multiprocessing.get_context('fork') Pool, Lock, Semaphore, Manager = _mp.Pool, _mp.Lock, _mp.Semaphore, _mp.Manager import hashlib import ctypes -from pymtp import MTP +from pymtp import LIBMTP_DeviceEntry, LIBMTP_RawDevice, MTP import string # --- per-board dev-session locks (see test/hil/board_lock.py) ------------ @@ -127,11 +128,11 @@ def enum_timeout() -> int: 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 - body (subprocess, libmtp scan) counts against the budget. Returns the first truthy - predicate value, or None on timeout.""" - deadline = time.monotonic() + enum_timeout() + body (subprocess, libmtp scan) counts against the budget. An explicit timeout overrides + that budget. Returns the first truthy predicate value, or None on timeout.""" + deadline = time.monotonic() + (enum_timeout() if timeout is None else timeout) while True: r = predicate() if r: @@ -503,23 +504,120 @@ def read_disk_file(uid: str, lun: int, fname: str) -> bytes: return data -def open_mtp_dev(uid): +def open_mtp_dev(uid: str): mtp = MTP() + last_usb = None + deadline = time.monotonic() + 2 * enum_timeout() - def try_open(): - # unmount gio/gvfs MTP mount which blocks libmtp from accessing the device - subprocess.run(f"gio mount -u mtp://TinyUsb_TinyUsb_Device_{uid}/", - shell=True, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL) - for raw in mtp.detect_devices(): - mtp.device = mtp.mtp.LIBMTP_Open_Raw_Device(ctypes.byref(raw)) - if mtp.device: - sn = mtp.get_serialnumber().decode('utf-8') - if sn == uid: - return mtp - mtp.disconnect() + def find_usb(): + nonlocal last_usb + for serial_fname in glob.glob('/sys/bus/usb/devices/*/serial'): + dev_path = Path(serial_fname).parent + try: + if (Path(serial_fname).read_text().strip().lower() != uid.lower() + or (dev_path / 'idVendor').read_text().strip() != 'cafe' + or (dev_path / 'idProduct').read_text().strip() != '4017'): + continue + busnum = int((dev_path / 'busnum').read_text()) + 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 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): @@ -1432,15 +1530,13 @@ def test_device_mtp(board): _null = os.open(os.devnull, os.O_WRONLY) os.dup2(_null, fd) - mtp = open_mtp_dev(uid) - - # --- AFTER: restore stderr --- - os.dup2(_saved, fd) - os.close(_null) - os.close(_saved) - - if mtp is None or mtp.device is None: - assert False, 'MTP device not found' + try: + mtp = open_mtp_dev(uid) + finally: + # --- AFTER: restore stderr --- + os.dup2(_saved, fd) + os.close(_null) + os.close(_saved) try: assert b"TinyUSB" == mtp.get_manufacturer(), 'MTP wrong manufacturer'