The biggest, the largest

This commit is contained in:
firestar5683
2026-08-11 18:52:27 -05:00
parent 4f471fc38c
commit dcde24f667
25 changed files with 396 additions and 492 deletions
@@ -158,6 +158,15 @@ def should_track_stop_accel_directly(stopping: bool, v_ego: float,
return bool(stopping and v_ego > EV6_GT_LINE_STOP_BRAKE_CAP_MAX_SPEED and accel_cmd < actual_accel)
def should_track_stop_accel_directly_for_car(car_fingerprint, stopping: bool, v_ego: float,
accel_cmd: float, actual_accel: float) -> bool:
# EV9 uses the low-speed stop cap to avoid a harsh lead re-brake handoff.
# Keep direct tracking for the other CCNC angle-steering platform.
return car_fingerprint != CAR.KIA_EV9 and should_track_stop_accel_directly(
stopping, v_ego, accel_cmd, actual_accel,
)
def should_use_ev6_gt_line_stop_direct_tracking(ev6_gt_line: bool, stopping: bool, v_ego: float,
accel_cmd: float, actual_accel: float) -> bool:
return bool(ev6_gt_line and stopping and v_ego > EV6_GT_LINE_STOP_BRAKE_CAP_MAX_SPEED and accel_cmd < actual_accel)
@@ -399,6 +408,11 @@ def process_hud_alert(enabled, fingerprint, hud_control):
return sys_warning, sys_state, left_lane_warning, right_lane_warning
def preserve_stock_canfd_lfa_status(car_fingerprint) -> bool:
# The 2022-24 Carnival expects a clean replacement status payload after its radar ECU is disabled.
return car_fingerprint != CAR.KIA_CARNIVAL_4TH_GEN
class CarController(CarControllerBase):
def __init__(self, dbc_names, CP):
super().__init__(dbc_names, CP)
@@ -630,8 +644,10 @@ class CarController(CarControllerBase):
if should_use_ev6_gt_line_stop_direct_tracking(is_ev6_gt_line, self._ioniq_6_long_tuning.stopping,
CS.out.vEgo, accel_cmd, self._ioniq_6_long_tuning.actual_accel):
use_egmp_smoothed_accel = False
if is_ccnc_angle_long and should_track_stop_accel_directly(self._ioniq_6_long_tuning.stopping, CS.out.vEgo,
accel_cmd, self._ioniq_6_long_tuning.actual_accel):
if is_ccnc_angle_long and should_track_stop_accel_directly_for_car(
self.CP.carFingerprint, self._ioniq_6_long_tuning.stopping, CS.out.vEgo,
accel_cmd, self._ioniq_6_long_tuning.actual_accel,
):
use_egmp_smoothed_accel = False
if use_egmp_dynamic_long_tuning:
if use_egmp_smoothed_accel:
@@ -803,10 +819,11 @@ class CarController(CarControllerBase):
forward_stock_lkas = angle_lkas_alt and (
angle_lkas_alt_standstill_handoff or not (drive_gear and (CC.latActive or CC.enabled))
)
preserve_stock_lfa_status = preserve_stock_canfd_lfa_status(self.CP.carFingerprint)
if not forward_stock_lkas and not ccnc_angle_long:
can_sends.extend(hyundaicanfd.create_steering_messages(self.packer, self.CP, self.CAN, CC.enabled,
steering_msg_active, apply_torque, apply_angle,
CS.stock_lfa_msg,
CS.stock_lfa_msg if preserve_stock_lfa_status else None,
CS.stock_lkas_msg if preserve_stock_lkas else None,
lka_icon=lka_icon,
send_lfa_status=self.ecu_disable_failed and
@@ -845,7 +862,8 @@ class CarController(CarControllerBase):
CC.leftBlinker, CC.rightBlinker, CS.msg_161, CS.msg_162, CS.msg_1b5,
CS.is_metric, CS.out, CS.out.cruiseState.available, lfa_icon))
else:
can_sends.append(hyundaicanfd.create_lfahda_cluster(self.packer, self.CAN, CC.enabled, CS.stock_lfahda_cluster_msg,
cluster_base_values = CS.stock_lfahda_cluster_msg if preserve_stock_lfa_status else None
can_sends.append(hyundaicanfd.create_lfahda_cluster(self.packer, self.CAN, CC.enabled, cluster_base_values,
lfa_icon=lfa_icon))
# blinkers
@@ -16,7 +16,9 @@ from opendbc.car.hyundai.carcontroller import CarController, Ioniq6LongitudinalT
get_canfd_scc_decel_step, \
should_reset_ev6_gt_line_longitudinal_tuning, reset_ev6_gt_line_longitudinal_tuning, \
direct_angle_request_allowed, get_angle_smoothing_alpha, \
should_use_ev6_gt_line_stop_direct_tracking
should_use_ev6_gt_line_stop_direct_tracking, \
should_track_stop_accel_directly_for_car, \
preserve_stock_canfd_lfa_status
from opendbc.car.hyundai.carstate import CarState, decode_canfd_camera_lead, decode_ioniq_6_blindspot_radar_state, \
get_canfd_cruise_available
from opendbc.car.hyundai.interface import CarInterface, KIA_EV9_ACCEL_MAX
@@ -122,6 +124,29 @@ def get_test_toggles() -> SimpleNamespace:
class TestHyundaiFingerprint:
def test_carnival_2024_uses_clean_canfd_lfa_status(self):
assert not preserve_stock_canfd_lfa_status(CAR.KIA_CARNIVAL_4TH_GEN)
assert preserve_stock_canfd_lfa_status(CAR.KIA_CARNIVAL_2025)
assert preserve_stock_canfd_lfa_status(CAR.KIA_CARNIVAL_HEV_4TH_GEN)
assert preserve_stock_canfd_lfa_status(CAR.HYUNDAI_IONIQ_6)
CP = CarParams.new_message()
CP.carFingerprint = CAR.KIA_CARNIVAL_4TH_GEN
CP.flags = int(HyundaiFlags.CANFD | HyundaiFlags.RADAR_SCC)
CP.openpilotLongitudinalControl = True
packer = CANPacker(DBC[CP.carFingerprint][Bus.pt])
can_bus = CanBus(CP)
stock_lfa = {"HAS_LANE_SAFETY": 1, "NEW_SIGNAL_4": 8, "DAMP_FACTOR": 100}
lfa_base = stock_lfa if preserve_stock_canfd_lfa_status(CP.carFingerprint) else None
lfa_msg = hyundaicanfd.create_steering_messages(packer, CP, can_bus, False, False, 0, 0.0, lfa_base)[0]
assert lfa_msg[1] == bytes.fromhex("05100002400008000000000000640000")
stock_cluster = {"NEW_SIGNAL_5": 1}
cluster_base = stock_cluster if preserve_stock_canfd_lfa_status(CP.carFingerprint) else None
cluster_msg = hyundaicanfd.create_lfahda_cluster(packer, can_bus, False, cluster_base)
assert cluster_msg[1] == bytes.fromhex("8e040000000000000000000000000000")
def test_canfd_torque_bsm_parser_registers_rear_blindspots(self):
CP = CarParams.new_message()
CP.carFingerprint = CAR.GENESIS_GV70_ELECTRIFIED_1ST_GEN
@@ -1417,6 +1442,14 @@ class TestHyundaiFingerprint:
assert get_canfd_scc_decel_step(ev9_cp) == pytest.approx(0.20)
assert get_canfd_scc_decel_step(ioniq_6_cp) == pytest.approx(0.36)
def test_ev9_keeps_low_speed_stop_brake_cap(self):
assert not should_track_stop_accel_directly_for_car(
CAR.KIA_EV9, stopping=True, v_ego=1.8, accel_cmd=-3.5, actual_accel=-0.4,
)
assert should_track_stop_accel_directly_for_car(
CAR.HYUNDAI_IONIQ_5_PE, stopping=True, v_ego=1.8, accel_cmd=-3.5, actual_accel=-0.4,
)
def test_ev9_longitudinal_decel_jerk_is_bounded(self):
state = Ioniq6LongitudinalTuningState(actual_accel=-1.0, accel_last=-1.0)
state = update_ioniq_6_longitudinal_tuning(
@@ -160,7 +160,7 @@ class CarController(CarControllerBase):
self.CP.openpilotLongitudinalControl, CC.longActive, hud_control.leadVisible,
self.status_bus))
can_sends.append(subarucan.create_es_lkas_state(self.packer, self.frame // 10, CS.es_lkas_state_msg, CC.enabled, hud_control.visualAlert,
can_sends.append(subarucan.create_es_lkas_state(self.packer, self.frame // 10, CS.es_lkas_state_msg, CC.latActive, hud_control.visualAlert,
hud_control.leftLaneVisible, hud_control.rightLaneVisible,
hud_control.leftLaneDepart, hud_control.rightLaneDepart, self.status_bus))
@@ -124,6 +124,7 @@ def create_es_lkas_state(packer, frame, es_lkas_state_msg, enabled, visual_alert
values["LKAS_ACTIVE"] = 1 # Show LKAS lane lines
values["LKAS_Dash_State"] = 2 # Green enabled indicator
else:
values["LKAS_ACTIVE"] = 0
values["LKAS_Dash_State"] = 0 # LKAS Not enabled
values["LKAS_Left_Line_Visible"] = int(left_line)
@@ -1,14 +1,18 @@
import inspect
from collections import defaultdict
from types import SimpleNamespace
import pytest
from opendbc.can import CANPacker, CANParser
from opendbc.car import Bus
from opendbc.car.subaru import subarucan
from opendbc.car.subaru.carcontroller import CarController
from opendbc.car.subaru.carstate import CarState
from opendbc.car.subaru.fingerprints import FW_VERSIONS
from opendbc.car.fw_versions import match_fw_to_car
from opendbc.car.subaru.interface import CarInterface
from opendbc.car.subaru.values import CAR, CanBus, SubaruFlags, SubaruSafetyFlags
from opendbc.car.subaru.values import CAR, DBC, CanBus, SubaruFlags, SubaruSafetyFlags
from opendbc.car.structs import CarParams
@@ -90,7 +94,6 @@ class TestSubaruFingerprint:
assert matches == {CAR.SUBARU_OUTBACK_2023}
def test_legacy_2025_firmware(self):
legacy_fw = FW_VERSIONS[CAR.SUBARU_LEGACY_2025]
car_fw = [
CarParams.CarFw(ecu=CarParams.Ecu.abs, fwVersion=b'\xa1 $\x11\x00', address=0x7b0, brand="subaru"),
CarParams.CarFw(ecu=CarParams.Ecu.eps, fwVersion=b'[\xc0\xd1\x10\x00', address=0x746, brand="subaru"),
@@ -172,23 +175,23 @@ def test_outback_2023_uses_d_platform_bus_layout():
assert controller.status_bus == CanBus.main
def test_legacy_2025_uses_d_platform_bus_layout():
def test_legacy_2025_uses_gen2_angle_bus_layout():
CP = CarInterface.get_non_essential_params(CAR.SUBARU_LEGACY_2025)
parsers = CarState.get_can_parsers(CP)
controller = CarController({}, CP)
assert CP.flags & SubaruFlags.D_PLATFORM
assert CP.safetyConfigs[0].safetyParam & SubaruSafetyFlags.D_PLATFORM
assert CP.flags & SubaruFlags.D_PLATFORM_CAMERA
assert CP.safetyConfigs[0].safetyParam & SubaruSafetyFlags.D_PLATFORM_CAMERA
assert CanBus.main_for_cp(CP) == CanBus.alt
assert CanBus.angle_for_cp(CP) == CanBus.camera
assert parsers[Bus.pt].bus == CanBus.alt
assert not (CP.flags & SubaruFlags.D_PLATFORM)
assert not (CP.safetyConfigs[0].safetyParam & SubaruSafetyFlags.D_PLATFORM)
assert not (CP.flags & SubaruFlags.D_PLATFORM_CAMERA)
assert not (CP.safetyConfigs[0].safetyParam & SubaruSafetyFlags.D_PLATFORM_CAMERA)
assert CanBus.main_for_cp(CP) == CanBus.main
assert CanBus.angle_for_cp(CP) == CanBus.main
assert parsers[Bus.pt].bus == CanBus.main
assert parsers[Bus.cam].bus == CanBus.camera
assert parsers[Bus.alt].bus == CanBus.alt
assert parsers[Bus.main].bus == CanBus.main
assert controller.angle_bus == CanBus.camera
assert controller.status_bus == CanBus.camera
assert Bus.main not in parsers
assert controller.angle_bus == CanBus.main
assert controller.status_bus == CanBus.main
def test_ascent_2023_uses_d_platform_bus_layout():
@@ -232,3 +235,26 @@ def test_angle_controller_tracks_driver_override():
assert controller.driver_override
assert controller.apply_steer_last == CS.out.steeringAngleDeg
assert msg[0] == 0x124
def test_lkas_hud_state_uses_lateral_active():
update_source = inspect.getsource(CarController.update)
assert "create_es_lkas_state(self.packer, self.frame // 10, CS.es_lkas_state_msg, CC.latActive" in update_source
assert "create_es_lkas_state(self.packer, self.frame // 10, CS.es_lkas_state_msg, CC.enabled" not in update_source
@pytest.mark.parametrize(("enabled", "expected"), ((False, 0), (True, 1)))
def test_lkas_hud_active_bit_follows_lateral_state(enabled, expected):
dbc = DBC[CAR.SUBARU_LEGACY_2025][Bus.pt]
packer = CANPacker(dbc)
parser = CANParser(dbc, [("ES_LKAS_State", 0)], CanBus.main)
stock_lkas_state = defaultdict(int, {"LKAS_ACTIVE": 1})
msg = subarucan.create_es_lkas_state(
packer, 0, stock_lkas_state, enabled, 0, False, False, False, False, CanBus.main,
)
parser.update([(1, [msg])])
assert parser.can_valid
assert parser.vl["ES_LKAS_State"]["LKAS_ACTIVE"] == expected
+1 -1
View File
@@ -242,7 +242,7 @@ class CAR(Platforms):
SUBARU_LEGACY_2025 = SubaruGen2PlatformConfig(
[SubaruCarDocs("Subaru Legacy 2025", "All", car_parts=CarParts.common([CarHarness.subaru_d]))],
SUBARU_OUTBACK.specs,
flags=SubaruFlags.LKAS_ANGLE | SubaruFlags.D_PLATFORM | SubaruFlags.D_PLATFORM_CAMERA,
flags=SubaruFlags.LKAS_ANGLE,
)
SUBARU_ASCENT_2023 = SubaruGen2PlatformConfig(
[SubaruCarDocs("Subaru Ascent 2023-25", "All", car_parts=CarParts.common([CarHarness.subaru_d]))],
@@ -90,6 +90,7 @@
{.msg = {{MSG_SUBARU_Wheel_Speeds, alt_bus, 8, 50U, .max_counter = 15U, .ignore_quality_flag = true}, { 0 }, { 0 }}}, \
{.msg = {{MSG_SUBARU_Brake_Status, alt_bus, 8, 50U, .max_counter = 15U, .ignore_quality_flag = true}, { 0 }, { 0 }}}, \
{.msg = {{MSG_SUBARU_ES_Status, status_bus, 8, 20U, .max_counter = 15U, .ignore_quality_flag = true}, { 0 }, { 0 }}}, \
{.msg = {{MSG_SUBARU_ES_DashStatus, SUBARU_CAM_BUS, 8, 10U, .max_counter = 15U, .ignore_quality_flag = true}, { 0 }, { 0 }}}, \
#define SUBARU_D_PLATFORM_ANGLE_RX_CHECKS() \
{.msg = {{MSG_SUBARU_Throttle, SUBARU_ALT_BUS, 8, 100U, .max_counter = 15U, .ignore_quality_flag = true}, { 0 }, { 0 }}}, \
@@ -148,6 +149,10 @@ static void subaru_rx_hook(const CANPacket_t *msg) {
pcm_cruise_check(cruise_engaged);
}
if (subaru_lkas_angle && (msg->addr == MSG_SUBARU_ES_DashStatus) && (msg->bus == SUBARU_CAM_BUS)) {
acc_main_on = GET_BIT(msg, 49U);
}
if (!subaru_lkas_angle && (msg->addr == MSG_SUBARU_CruiseControl) && (msg->bus == alt_main_bus)) {
bool cruise_engaged = (msg->data[5] >> 1) & 1U;
pcm_cruise_check(cruise_engaged);
@@ -9,6 +9,7 @@ from opendbc.car.subaru.carcontroller import get_safety_CP
from opendbc.car.subaru.values import CarControllerParams, SubaruSafetyFlags
from opendbc.car.structs import CarParams
from opendbc.car.vehicle_model import VehicleModel
from opendbc.safety import ALTERNATIVE_EXPERIENCE
from opendbc.safety.tests.libsafety import libsafety_py
import opendbc.safety.tests.common as common
from opendbc.safety.tests.common import CANPackerSafety, away_round, round_speed
@@ -230,6 +231,21 @@ class TestSubaruAngleSafetyBase(TestSubaruSafetyBase, common.AngleSteeringSafety
def _toggle_aol(self, toggle_on):
return None
def _acc_main_msg(self, main_on):
values = {"Cruise_On": int(main_on)}
return self.packer.make_can_msg_panda("ES_DashStatus", SUBARU_CAM_BUS, values)
def test_acc_main_tracks_dash_status(self):
self.safety.set_alternative_experience(ALTERNATIVE_EXPERIENCE.ALWAYS_ON_LATERAL)
self._rx(self._acc_main_msg(False))
self.assertFalse(self.safety.get_acc_main_on())
self.assertFalse(self.safety.get_aol_allowed())
self._rx(self._acc_main_msg(True))
self.assertTrue(self.safety.get_acc_main_on())
self.assertTrue(self.safety.get_aol_allowed())
def test_angle_cmd_when_enabled(self):
pass
@@ -558,6 +558,9 @@ class LatControlTorque(LatControl):
elif genesis_g70_active:
output_torque *= genesis_g70_center_output_taper
output_torque *= get_genesis_g70_curve_unwind_output_scale(setpoint, desired_lateral_jerk, CS.vEgo)
output_torque *= get_genesis_g70_high_speed_error_scale(
setpoint, measurement, desired_lateral_jerk, CS.vEgo,
)
low_speed_output_limit = get_genesis_g70_low_speed_output_limit(setpoint, CS.vEgo)
output_torque = float(np.clip(output_torque, -low_speed_output_limit, low_speed_output_limit))
elif sonata_hybrid_active:
@@ -252,6 +252,13 @@ GENESIS_G70_CURVE_UNWIND_LAT = 0.25
GENESIS_G70_CURVE_UNWIND_LAT_WIDTH = 0.12
GENESIS_G70_CURVE_UNWIND_JERK = 0.08
GENESIS_G70_CURVE_UNWIND_JERK_WIDTH = 0.08
GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_MAX = 0.15
GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_SPEED = 50.0 * CV.MPH_TO_MS
GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_SPEED_WIDTH = 8.0 * CV.MPH_TO_MS
GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_ERROR = 0.18
GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_ERROR_WIDTH = 0.15
GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_JERK = 0.15
GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_JERK_WIDTH = 0.10
BOLT_2017_LATERAL_TESTING_GROUND_ID = testing_ground.id_3
BOLT_2017_STEER_RATIO_TEST_SCALE = 1.045
@@ -321,12 +328,14 @@ BOLT_2022_2023_CENTER_TAPER_LAT = 0.18
BOLT_2022_2023_CENTER_TAPER_LAT_WIDTH = 0.03
BOLT_2022_2023_CENTER_TAPER_SPEED = 25.0
BOLT_2022_2023_CENTER_TAPER_SPEED_WIDTH = 2.5
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_MAX = 0.07
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_MAX = 0.12
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_LAT = 0.14
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_LAT_WIDTH = 0.04
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_SPEED = 4.0
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_SPEED_WIDTH = 1.5
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_SPEED_MAX = 14.0
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_FLOOR = 2.0
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_FLOOR_WIDTH = 0.7
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_SPEED_MAX = 16.5
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_SPEED_MAX_WIDTH = 2.0
BOLT_2022_2023_LOW_SPEED_CENTER_OUTPUT_LIMIT = 0.38
BOLT_2022_2023_LOW_SPEED_CENTER_OUTPUT_LAT = 0.17
@@ -1903,10 +1912,14 @@ def get_bolt_2022_2023_center_output_scale(desired_lateral_accel: float, v_ego:
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_SPEED_MAX_WIDTH)
low_speed_center_weight = _bolt_2022_2023_sigmoid((BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_LAT - abs(desired_lateral_accel)) /
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_LAT_WIDTH)
low_speed_floor = _bolt_2022_2023_sigmoid(
(v_ego - BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_FLOOR) /
BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_FLOOR_WIDTH
)
highway_reduction = (_flm_vehicle_knob("gm_bolt_2022_2023.center_taper_max", BOLT_2022_2023_CENTER_TAPER_MAX) *
highway_speed_weight * highway_center_weight)
low_speed_reduction = (BOLT_2022_2023_LOW_SPEED_CENTER_TAPER_MAX * low_speed_onset * low_speed_cutoff *
low_speed_center_weight)
low_speed_center_weight * low_speed_floor)
return 1.0 - min(highway_reduction + low_speed_reduction, 0.95)
@@ -2685,6 +2698,23 @@ def get_genesis_g70_curve_unwind_output_scale(desired_lateral_accel: float, desi
return 1.0 + GENESIS_G70_CURVE_UNWIND_OUTPUT_BOOST * speed_weight * lateral_weight * jerk_weight
def get_genesis_g70_high_speed_error_scale(setpoint: float, measured_lateral_accel: float,
desired_lateral_jerk: float, v_ego: float) -> float:
tracking_error = abs(measured_lateral_accel - setpoint)
if tracking_error <= 0.0:
return 1.0
speed_weight = _sigmoid((v_ego - GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_SPEED) /
GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_SPEED_WIDTH)
error_weight = _sigmoid((tracking_error - GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_ERROR) /
GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_ERROR_WIDTH)
jerk_weight = _sigmoid((abs(desired_lateral_jerk) - GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_JERK) /
GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_JERK_WIDTH)
phase_weight = 1.0 if setpoint * desired_lateral_jerk < 0.0 else 0.45
reduction = (GENESIS_G70_HIGH_SPEED_ERROR_DAMPING_MAX * speed_weight * error_weight *
(0.35 + (0.65 * jerk_weight)) * phase_weight)
return 1.0 - reduction
def _ioniq_5_sigmoid(x: float) -> float:
return _sigmoid(x)
@@ -4,6 +4,7 @@ import numpy as np
HONDA_HRV_3G_FAR_FOLLOW_BRAKE_SLEW_RATE = 3.0
HONDA_HRV_3G_FAR_FOLLOW_RELEASE_SLEW_RATE = 2.0
HONDA_HRV_3G_UNTRACKED_SLOW_LEAD_DECEL_SCALE = 1.35
HYUNDAI_ELANTRA_LEAD_FOLLOW_JERK_SCALE = 1.25
GM_SILVERADO_EARLY_FOLLOW_MIN_EGO_SPEED = 18.0
GM_SILVERADO_EARLY_FOLLOW_MAX_DISTANCE = 130.0
GM_SILVERADO_EARLY_FOLLOW_MIN_MODEL_PROB = 0.85
@@ -42,6 +43,13 @@ def get_untracked_slow_lead_decel_scale(CP):
return 1.0
def get_lead_follow_jerk_scale(CP):
"""Spread the lead-source transition for cars with a sharp vision-lead handoff."""
if getattr(CP, "brand", "") == "hyundai" and str(getattr(CP, "carFingerprint", "")) == "HYUNDAI_ELANTRA_2021":
return HYUNDAI_ELANTRA_LEAD_FOLLOW_JERK_SCALE
return 1.0
def is_gm_silverado_early_follow_lead(CP, lead, v_ego):
"""Admit a credible centered vision lead before it becomes a close lead."""
if (
+7 -2
View File
@@ -74,6 +74,7 @@ from openpilot.selfdrive.controls.lib.latcontrol_torque import (
get_genesis_g70_curve_unwind_output_scale,
get_genesis_g70_friction_jerk_deadzone,
get_genesis_g70_friction_threshold,
get_genesis_g70_high_speed_error_scale,
get_genesis_g70_low_speed_angle_damping,
get_genesis_g70_low_speed_output_limit,
get_genesis_gv70_friction_threshold,
@@ -286,9 +287,9 @@ class TestLatControl:
highway_turn = get_bolt_2022_2023_center_output_scale(0.40, 31.0)
creep_center = get_bolt_2022_2023_center_output_scale(0.04, 1.0)
assert 0.92 < low_speed_center < 0.95
assert 0.88 < low_speed_center < 0.92
assert low_speed_turn > 0.99
assert middle_speed_center > 0.98
assert 0.96 < middle_speed_center < 0.98
assert 0.88 < highway_center < 0.91
assert highway_turn > 0.99
assert creep_center > 0.99
@@ -812,6 +813,10 @@ class TestLatControl:
assert get_genesis_g70_curve_unwind_output_scale(0.7, -0.5, 25.0) > 1.0
assert get_genesis_g70_curve_unwind_output_scale(0.7, 0.5, 25.0) == 1.0
assert get_genesis_g70_high_speed_error_scale(0.2, 0.2, 0.8, 20.0) == 1.0
assert get_genesis_g70_high_speed_error_scale(0.2, 0.9, 0.8, 20.0) < 1.0
assert get_genesis_g70_high_speed_error_scale(0.2, 0.9, 0.8, 10.0) > get_genesis_g70_high_speed_error_scale(0.2, 0.9, 0.8, 20.0)
def test_sonata_hybrid_center_output_taper_is_mid_speed_and_center_gated(self):
low_speed = get_sonata_hybrid_center_output_scale(0.0, 8.0)
center = get_sonata_hybrid_center_output_scale(0.0, 13.4)
@@ -4,6 +4,7 @@ from types import SimpleNamespace
from openpilot.common.realtime import DT_MDL
from openpilot.starpilot.controls.starpilot_planner import StarPilotPlanner, get_force_stop_jerk_scale
from openpilot.selfdrive.controls.lib.longitudinal_vehicle_tunes import get_lead_follow_jerk_scale
import openpilot.starpilot.controls.starpilot_planner as starpilot_planner_module
@@ -38,6 +39,11 @@ def test_force_stop_jerk_scale_is_platform_specific():
assert get_force_stop_jerk_scale(SimpleNamespace(carFingerprint="OTHER_CAR")) == 0.32
def test_lead_follow_jerk_scale_is_platform_specific():
assert get_lead_follow_jerk_scale(SimpleNamespace(brand="hyundai", carFingerprint="HYUNDAI_ELANTRA_2021")) == 1.25
assert get_lead_follow_jerk_scale(SimpleNamespace(brand="other", carFingerprint="OTHER_CAR")) == 1.0
def make_sm(planner, *, frame: int, v_ego: float, left_blinker: bool, right_blinker: bool = False, standstill: bool = False):
return FakeSM(frame, {
"radarState": SimpleNamespace(
+32 -19
View File
@@ -1,42 +1,55 @@
#!/usr/bin/env python3
import sys
import pathlib
import codecs
import pickle
import pathlib
import sys
from typing import Any
import onnx
from tinygrad.nn.onnx import OnnxPBParser
def get_name_and_shape(value_info:onnx.ValueInfoProto) -> tuple[str, tuple[int,...]]:
shape = tuple([int(dim.dim_value) for dim in value_info.type.tensor_type.shape.dim])
name = value_info.name
class MetadataOnnxPBParser(OnnxPBParser):
def _parse_ModelProto(self) -> dict:
obj: dict[str, Any] = {"graph": {"input": [], "output": []}, "metadata_props": []}
for fid, wire_type in self._parse_message(self.reader.len):
match fid:
case 7:
obj["graph"] = self._parse_GraphProto()
case 14:
obj["metadata_props"].append(self._parse_StringStringEntryProto())
case _:
self.reader.skip_field(wire_type)
return obj
def get_name_and_shape(value_info: dict[str, Any]) -> tuple[str, tuple[int, ...]]:
shape = tuple(int(dim) if isinstance(dim, int) else 0 for dim in value_info["parsed_type"].shape)
name = value_info["name"]
return name, shape
def get_metadata_value_by_name(model:onnx.ModelProto, name:str) -> str | Any:
for prop in model.metadata_props:
if prop.key == name:
return prop.value
def get_metadata_value_by_name(model: dict[str, Any], name: str) -> str | Any:
for prop in model["metadata_props"]:
if prop["key"] == name:
return prop["value"]
return None
def make_metadata_dict(model_path: str | pathlib.Path) -> dict[str, Any]:
model = onnx.load(str(model_path))
def make_metadata_dict(model_path):
model = MetadataOnnxPBParser(model_path).parse()
output_slices = get_metadata_value_by_name(model, 'output_slices')
assert output_slices is not None, 'output_slices not found in metadata'
return {
'model_checkpoint': get_metadata_value_by_name(model, 'model_checkpoint'),
'output_slices': pickle.loads(codecs.decode(output_slices.encode(), "base64")),
'input_shapes': dict([get_name_and_shape(x) for x in model.graph.input]),
'output_shapes': dict([get_name_and_shape(x) for x in model.graph.output]),
'input_shapes': dict(get_name_and_shape(x) for x in model["graph"]["input"]),
'output_shapes': dict(get_name_and_shape(x) for x in model["graph"]["output"]),
}
if __name__ == "__main__":
model_path = pathlib.Path(sys.argv[1])
metadata = make_metadata_dict(model_path)
metadata_path = model_path.parent / (model_path.stem + '_metadata.pkl')
with open(metadata_path, 'wb') as f:
pickle.dump(metadata, f)
pickle.dump(make_metadata_dict(model_path), f)
print(f'saved metadata to {metadata_path}')
@@ -0,0 +1,35 @@
import base64
import pickle
import onnx
from openpilot.selfdrive.modeld.get_model_metadata import make_metadata_dict
def test_make_metadata_dict_uses_disk_backed_parser(tmp_path):
model_input = onnx.helper.make_tensor_value_info("input", onnx.TensorProto.FLOAT, (1, 3))
model_output = onnx.helper.make_tensor_value_info("output", onnx.TensorProto.FLOAT, (1, 4))
graph = onnx.helper.make_graph(
[onnx.helper.make_node("Identity", ["input"], ["output"])],
"metadata-test",
[model_input],
[model_output],
)
model = onnx.helper.make_model(graph)
output_slices = {"path": slice(0, 4)}
model.metadata_props.append(onnx.StringStringEntryProto(
key="output_slices",
value=base64.b64encode(pickle.dumps(output_slices)).decode(),
))
model.metadata_props.append(onnx.StringStringEntryProto(key="model_checkpoint", value="test-checkpoint"))
model_path = tmp_path / "model.onnx"
onnx.save(model, model_path)
metadata = make_metadata_dict(model_path)
assert metadata == {
"model_checkpoint": "test-checkpoint",
"output_slices": output_slices,
"input_shapes": {"input": (1, 3)},
"output_shapes": {"output": (1, 4)},
}
@@ -119,6 +119,13 @@ class DeveloperLayoutMici(NavScroller):
super()._update_state()
self._ssh_fetcher.update()
def show_event(self):
super().show_event()
# CarParamsPersistent can change while this panel is not visible (for
# example after applying a manual fingerprint). Refresh capability-gated
# controls when the panel is opened so stale controls are not shown.
self._update_toggles()
def _update_toggles(self):
ui_state.update_params()
@@ -66,7 +66,7 @@ def test_external_gpu_requirement_is_cached_from_manifest(tmp_path, monkeypatch)
def test_external_gpu_compilation_is_opt_in(tmp_path, monkeypatch):
invocations = []
monkeypatch.setattr(model_compiler, "build_compile_env", lambda: {
monkeypatch.setattr(model_compiler, "build_compile_env", lambda **_: {
"DEV": "QCOM", "IMAGE": "2", "NOLOCALS": "1", "OPENPILOT_HACKS": "1",
})
monkeypatch.setattr(model_compiler.subprocess, "run", lambda command, **kwargs: invocations.append((command, kwargs)))
@@ -88,7 +88,7 @@ def test_external_gpu_compilation_is_opt_in(tmp_path, monkeypatch):
def test_compile_clears_only_selected_model_outputs(tmp_path, monkeypatch):
monkeypatch.setattr(model_compiler, "build_compile_env", lambda: {})
monkeypatch.setattr(model_compiler, "build_compile_env", lambda **_: {})
monkeypatch.setattr(model_compiler.subprocess, "run", lambda *args, **kwargs: None)
(tmp_path / "normal_driving_tinygrad.pkl").write_bytes(b"old")
(tmp_path / "normal_driving_tinygrad.pkl.p00").write_bytes(b"old")
+12 -4
View File
@@ -18,6 +18,7 @@ from openpilot.selfdrive.controls.lib.lead_behavior import (
should_track_lead,
)
from openpilot.selfdrive.controls.lib.longitudinal_mpc_lib.long_mpc import A_CHANGE_COST, DANGER_ZONE_COST, J_EGO_COST, STOP_DISTANCE
from openpilot.selfdrive.controls.lib.longitudinal_vehicle_tunes import get_lead_follow_jerk_scale
from openpilot.starpilot.common.starpilot_utilities import calculate_lane_width, calculate_road_curvature
from openpilot.starpilot.common.starpilot_variables import CRUISING_SPEED, MINIMUM_LATERAL_ACCELERATION, PLANNER_TIME, THRESHOLD
@@ -300,12 +301,19 @@ class StarPilotPlanner:
# While committed to a Force Stop, cut the MPC's accel-change penalty so terminal
# braking can ramp faster. 0.32 lands near 40, what long_mpc uses in blended mode.
try:
car_params = sm["carParams"]
except (KeyError, IndexError, TypeError, AttributeError):
car_params = None
if self.starpilot_vcruise.forcing_stop:
try:
car_params = sm["carParams"]
except (KeyError, IndexError, TypeError, AttributeError):
car_params = None
jerk_scale = get_force_stop_jerk_scale(car_params)
elif self.tracking_lead:
# Elantra vision leads can hand off from cruise to lead0 while closing
# quickly. A slightly higher accel-change cost makes that handoff begin
# earlier instead of arriving as a sharp brake request, without changing
# the safety stop distance or the force-stop path.
jerk_scale = get_lead_follow_jerk_scale(car_params)
else:
jerk_scale = 1.0
starpilotPlan.accelerationJerk = float(A_CHANGE_COST * self.starpilot_following.acceleration_jerk * jerk_scale)
+4 -3
View File
@@ -8,9 +8,10 @@ Quick start:
* set `STRICT_MODE=1` to kill the app if it drops too much below 60fps
* set `SCALE=1.5` to scale the entire UI by 1.5x
* set `BURN_IN=1` to get a burn-in heatmap version of the UI
* burn-in prevention shifts the UI by 2 pixels every 3 minutes on device; set `BURN_IN_PREVENTION=0` to disable it
or tune it with `BURN_IN_SHIFT_PIXELS` and `BURN_IN_SHIFT_INTERVAL` (seconds). TICI/TIZI and MICI use direct shifting
by default; setting `WHITE_LUMINANCE_CAP` below `1.0` enables the optional luminance cap and offscreen presentation path.
* burn-in prevention shifts the UI by 2 pixels every 3 minutes on device; the shift is blended over 1 second so it is not visible as a jump.
Set `BURN_IN_PREVENTION=0` to disable it or tune it with `BURN_IN_SHIFT_PIXELS`, `BURN_IN_SHIFT_INTERVAL`, and
`BURN_IN_SHIFT_TRANSITION_SECONDS`. TICI/TIZI shift the completed frame through the offscreen presentation path;
MICI uses direct shifting by default. Setting `WHITE_LUMINANCE_CAP` below `1.0` also enables the offscreen presentation path.
* set `MICI_FORCE_RENDER_TEXTURE=1` to force the C4 UI through the offscreen presentation path for diagnostics
* set `GRID=50` to show a 50-pixel alignment grid overlay
* set `MAGIC_DEBUG=1` to show every dropped frames (only on device)
+23 -4
View File
@@ -42,6 +42,10 @@ MICI_FORCE_RENDER_TEXTURE = os.getenv("MICI_FORCE_RENDER_TEXTURE", "0") == "1"
BURN_IN_PREVENTION = os.getenv("BURN_IN_PREVENTION", "0" if PC else "1") == "1"
BURN_IN_SHIFT_INTERVAL = max(1.0, float(os.getenv("BURN_IN_SHIFT_INTERVAL", "180")))
BURN_IN_SHIFT_PIXELS = max(0, int(os.getenv("BURN_IN_SHIFT_PIXELS", "2")))
BURN_IN_SHIFT_TRANSITION_SECONDS = min(
BURN_IN_SHIFT_INTERVAL,
max(0.1, float(os.getenv("BURN_IN_SHIFT_TRANSITION_SECONDS", "1"))),
)
WHITE_LUMINANCE_CAP = min(1.0, max(0.0, float(os.getenv(
"WHITE_LUMINANCE_CAP", "1.0"
))))
@@ -616,8 +620,11 @@ class GuiApplication:
self._render_texture_width = max(1, int(round(self._scaled_width * self._pixel_scale_x)))
self._render_texture_height = max(1, int(round(self._scaled_height * self._pixel_scale_y)))
# Keep raybig burn-in movement in final-frame composition. Translating the live EGL
# camera/widget pass can corrupt the camera presentation instead of shifting the UI.
needs_render_texture = ((self._scale != 1.0 and not PC) or BURN_IN_MODE or RECORD or
MICI_FORCE_RENDER_TEXTURE or
(BURN_IN_PREVENTION and DEVICE_TYPE != "mici") or
WHITE_LUMINANCE_CAP < 1.0)
if PC and self._scale != 1.0:
rl.set_mouse_scale(1 / self._scale, 1 / self._scale)
@@ -1138,13 +1145,25 @@ class GuiApplication:
except KeyboardInterrupt:
pass
def _burn_in_shift(self, now: float | None = None) -> tuple[int, int]:
def _burn_in_shift(self, now: float | None = None) -> tuple[float, float]:
if not BURN_IN_PREVENTION or BURN_IN_SHIFT_PIXELS == 0:
return 0, 0
return 0.0, 0.0
elapsed = (time.monotonic() if now is None else now) - self._burn_in_start_time
pattern_index = int(max(0.0, elapsed) // BURN_IN_SHIFT_INTERVAL) % len(BURN_IN_SHIFT_PATTERN)
x, y = BURN_IN_SHIFT_PATTERN[pattern_index]
elapsed = max(0.0, elapsed)
pattern_count = len(BURN_IN_SHIFT_PATTERN)
cycle_elapsed = elapsed % (BURN_IN_SHIFT_INTERVAL * pattern_count)
pattern_index = int(cycle_elapsed // BURN_IN_SHIFT_INTERVAL)
segment_elapsed = cycle_elapsed - pattern_index * BURN_IN_SHIFT_INTERVAL
# Blend into the next position at the end of each interval. This keeps the
# burn-in protection active without teleporting the entire UI by two pixels.
transition_start = BURN_IN_SHIFT_INTERVAL - BURN_IN_SHIFT_TRANSITION_SECONDS
transition = min(1.0, max(0.0, (segment_elapsed - transition_start) / BURN_IN_SHIFT_TRANSITION_SECONDS))
start_x, start_y = BURN_IN_SHIFT_PATTERN[pattern_index]
end_x, end_y = BURN_IN_SHIFT_PATTERN[(pattern_index + 1) % pattern_count]
x = start_x + (end_x - start_x) * transition
y = start_y + (end_y - start_y) * transition
return x * BURN_IN_SHIFT_PIXELS, y * BURN_IN_SHIFT_PIXELS
def font(self, font_weight: FontWeight = FontWeight.NORMAL) -> rl.Font:
+16
View File
@@ -0,0 +1,16 @@
from openpilot.system.ui.lib import application
def test_burn_in_shift_transitions_between_positions(monkeypatch):
app = object.__new__(application.GuiApplication)
app._burn_in_start_time = 100.0
monkeypatch.setattr(application, "BURN_IN_PREVENTION", True)
monkeypatch.setattr(application, "BURN_IN_SHIFT_INTERVAL", 10.0)
monkeypatch.setattr(application, "BURN_IN_SHIFT_PIXELS", 2)
monkeypatch.setattr(application, "BURN_IN_SHIFT_TRANSITION_SECONDS", 2.0)
assert app._burn_in_shift(108.0) == (0.0, 0.0)
midpoint = app._burn_in_shift(109.0)
assert midpoint == (-1.0, 0.0)
assert app._burn_in_shift(110.0) == (-2.0, 0.0)
@@ -1,37 +0,0 @@
import pytest
from tinygrad.runtime.support import system
class FakeUSB:
def __init__(self, product: str, is_custom: bool):
self.product, self.is_custom = product, is_custom
def patch_usb_dependencies(monkeypatch, usb):
calls = []
monkeypatch.setattr(system.System, "flock_acquire", lambda _: object())
monkeypatch.setattr(system, "USB3", lambda *args, **kwargs: calls.append((args, kwargs)) or usb)
return calls
def test_usb_gpu_rejects_legacy_bridge_firmware(monkeypatch):
calls = patch_usb_dependencies(monkeypatch, FakeUSB("USB 3.2 PCIe TinyEnclosure", False))
with pytest.raises(RuntimeError, match="unsupported legacy USB GPU firmware"):
system.USBPCIDevice("AM", object(), "usb:3-6")
assert calls[0][1] == {"use_bot": True}
def test_usb_gpu_accepts_custom_bridge_firmware(monkeypatch):
usb = FakeUSB("custom ASM2464PD", True)
calls = patch_usb_dependencies(monkeypatch, usb)
controller = object()
monkeypatch.setattr(system, "CustomASM24Controller", lambda candidate: controller if candidate is usb else None)
monkeypatch.setattr(system.System, "pci_setup_usb_bars", lambda *args, **kwargs: {2: (0, 1)})
device = system.USBPCIDevice("AM", object(), "usb:3-6")
assert calls[0][1] == {"use_bot": True}
assert device.usb is controller
+4 -15
View File
@@ -1,6 +1,6 @@
from __future__ import annotations
from typing import cast
import os, ctypes, struct, hashlib, functools, importlib, mmap, errno, array, contextlib, sys, weakref, itertools, collections, atexit, time
import os, ctypes, struct, hashlib, functools, importlib, mmap, errno, array, contextlib, sys, weakref, itertools, collections, atexit
assert sys.platform != 'win32'
from dataclasses import dataclass
from tinygrad.runtime.support.hcq import HCQCompiled, HCQAllocator, HCQBuffer, HWQueue, CLikeArgsState, HCQSignal, HCQProgram, FileIOInterface
@@ -909,21 +909,15 @@ class PCIIface(PCIIfaceBase):
class USBIface(PCIIface):
def __init__(self, dev, dev_id): # pylint: disable=super-init-not-called
deadline, visible = time.monotonic() + 5.0, []
while dev_id >= len(visible) and time.monotonic() < deadline:
visible = hcq_filter_visible_devices(USB3.list_devices(0xADD1, 0x0001) + USB3.list_devices(0x3801, 0x0001), "AMD")
if dev_id >= len(visible): time.sleep(0.1)
if dev_id >= len(visible):
if dev_id >= len(visible:=hcq_filter_visible_devices(USB3.list_devices(0xADD1, 0x0001) + USB3.list_devices(0x3801, 0x0001), "AMD")):
raise RuntimeError(f"AMD:{dev_id} does not exist ({pluralize('device', len(visible))} available)")
self.dev, self.pci_dev, self.vram_bar, self.count = dev, USBPCIDevice("AM", *visible[dev_id]), 0, len(visible)
self.dev_impl = AMDev(self.pci_dev)
self._compute_props()
self.pci_dev.usb._pci_cacheable += [self.pci_dev.bar_info(2)] # doorbell region is cacheable
# special regions
self.copy_bufs = [self._dma_region(ctrl_addr=0xf000, sys_addr=0x200000, size=0x80000)]
sys_next_off = 0x200 if self.pci_dev.usb.usb.is_custom else 0x800
self.sys_buf, self.sys_next_off = self._dma_region(ctrl_addr=0xa000, sys_addr=0x820000, size=0x1000), sys_next_off
self.sys_buf, self.sys_next_off = self._dma_region(ctrl_addr=0xa000, sys_addr=0x820000, size=0x1000), 0x200
self.cq_buf = self._dma_region(ctrl_addr=0xb800, sys_addr=0x822000, size=0x1000)
def _dma_region(self, ctrl_addr, sys_addr, size):
@@ -932,18 +926,13 @@ class USBIface(PCIIface):
def alloc(self, size:int, host=False, uncached=False, cpu_access=False, contiguous=False, force_devmem=False, **kwargs) -> HCQBuffer:
# usb allocates uncached and cpu_access in vram. vram writes are faster than sram writes
if (host or (not self.pci_dev.usb.usb.is_custom and uncached and cpu_access)) and self.sys_next_off + size < self.sys_buf.size:
if host and self.sys_next_off + size < self.sys_buf.size:
self.sys_next_off += size
return self.sys_buf.offset(self.sys_next_off - size, size)
# force devmem
return super().alloc(size, host=False, uncached=uncached, cpu_access=cpu_access, contiguous=contiguous, force_devmem=True, **kwargs)
def create_queue(self, queue_type, ring, gart, rptr, wptr, eop_buffer=None, cwsr_buffer=None, ctl_stack_size=0, ctx_save_restore_size=0,
xcc_id=0, idx=0):
if queue_type == kfd.KFD_IOC_QUEUE_TYPE_COMPUTE: self.pci_dev.usb._pci_cacheable += [(ring.cpu_view().addr, ring.size)]
return super().create_queue(queue_type, ring, gart, rptr, wptr, eop_buffer, cwsr_buffer, ctl_stack_size, ctx_save_restore_size, xcc_id, idx)
def sleep(self, timeout): pass
def _mock(iface, name=None): return type(name or f"MOCK{iface.__name__}", (iface,), {})
@@ -225,10 +225,8 @@ class USBPCIDevice(PCIDevice):
def __init__(self, devpref:str, dev, pcibus):
self.pcibus, self.peer_group = pcibus, f"USBPCIDevice_{pcibus}"
self.lock_fd = System.flock_acquire(f"{devpref.lower()}_{pcibus.lower()}.lock")
usb = USB3(dev, 0x81, 0x83, 0x02, 0x04, use_bot=True)
usb = USB3(dev)
if DEBUG >= 1: print(f"am {self.pcibus}: product string: {usb.product!r}")
if not usb.is_custom:
raise RuntimeError(f"unsupported legacy USB GPU firmware ({usb.product!r}); flash the current tinygrad ASM2464PD custom firmware")
self.usb: CustomASM24Controller = CustomASM24Controller(usb)
self._bar_info = System.pci_setup_usb_bars(self.usb, gpu_bus=4, mem_base=0x10000000, pref_mem_base=(32 << 30))
self.sram = BumpAllocator(size=0x80000, wrap=False) # asm24 controller sram
+84 -380
View File
@@ -1,7 +1,6 @@
import ctypes, struct, dataclasses, array, itertools, time, functools
from typing import Sequence
import ctypes, struct, time, functools, itertools
from tinygrad.runtime.autogen import libusb
from tinygrad.helpers import DEBUG, DEV, to_mv, round_up, OSX, getenv, ceildiv
from tinygrad.helpers import DEBUG, DEV, to_mv, round_up, ceildiv
from tinygrad.runtime.support.hcq import MMIOInterface
from tinygrad.runtime.support import c
@@ -22,6 +21,7 @@ class USB3:
return ctx
@classmethod
@functools.cache
def list_devices(cls, vendor:int, dev:int) -> list[tuple[c.POINTER[libusb.struct_libusb_device], str]]:
ret = []
for i in range(checked(libusb.libusb_get_device_list)(cls.ctx(), devs:=ctypes.POINTER(ctypes.POINTER(libusb.struct_libusb_device))())):
@@ -31,31 +31,12 @@ class USB3:
libusb.libusb_free_device_list(devs, 1)
return ret
@classmethod
def reopen_device(cls, vendor:int, product:int, bus_number:int, timeout:float=10.0):
deadline = time.monotonic() + timeout
while True:
for dev, _ in cls.list_devices(vendor, product):
handle = ctypes.POINTER(libusb.struct_libusb_device_handle)()
rc = libusb.libusb_open(dev, ctypes.byref(handle)) if libusb.libusb_get_bus_number(dev) == bus_number else libusb.LIBUSB_ERROR_NOT_FOUND
libusb.libusb_unref_device(dev)
if rc == 0:
return handle
if time.monotonic() >= deadline:
raise RuntimeError(f"device {vendor:04x}:{product:04x} did not reappear after reset")
time.sleep(0.1)
def __init__(self, dev:c.POINTER[libusb.struct_libusb_device], *args, **kwargs):
self._tags, self._transferred = itertools.count(1), ctypes.c_int(0)
self._bulk_buf, self._bulk_mv = alloc_cbuffer(4 << 20)
self._ctrl_buf, self._ctrl_mv = alloc_cbuffer(0x1000)
def __init__(self, dev:c.POINTER[libusb.struct_libusb_device], ep_data_in:int, ep_stat_in:int, ep_data_out:int, ep_cmd_out:int,
max_streams:int=31, use_bot=False):
self.ep_data_in, self.ep_stat_in, self.ep_data_out, self.ep_cmd_out = ep_data_in, ep_stat_in, ep_data_out, ep_cmd_out
self.max_streams, self.use_bot = max_streams, use_bot
self._transferred = ctypes.c_int(0)
self._bulk_in_buf, self._bulk_in_mv = alloc_cbuffer(4 << 20)
self._bulk_out_buf, self._bulk_out_mv = alloc_cbuffer(4 << 20)
bus_number = libusb.libusb_get_bus_number(dev)
self.handle = c.init_c_var(ctypes.POINTER(libusb.struct_libusb_device_handle), lambda x: checked(libusb.libusb_open)(dev, x))
libusb.libusb_unref_device(dev)
self.handle = c.init_c_var(c.POINTER[libusb.struct_libusb_device_handle], lambda x: checked(libusb.libusb_open)(dev, x))
# Read product string descriptor
_buf = (ctypes.c_ubyte * 256)()
@@ -63,210 +44,83 @@ class USB3:
checked(libusb.libusb_get_device_descriptor)(libusb.libusb_get_device(self.handle), ctypes.byref(_desc))
_ret = checked(libusb.libusb_get_string_descriptor_ascii)(self.handle, _desc.iProduct, _buf, 256)
self.product = bytes(_buf[:_ret]).decode("ascii", errors="replace")
self.is_custom = self.product.startswith("custom")
if self.is_custom: self.use_bot = use_bot = True
assert self.product.startswith("custom") or self.product.startswith("AS2462")
# Detach kernel driver if needed
if checked(libusb.libusb_kernel_driver_active)(self.handle, 0):
checked(libusb.libusb_detach_kernel_driver)(self.handle, 0)
reset_rc = libusb.libusb_reset_device(self.handle)
if reset_rc in (libusb.LIBUSB_ERROR_NO_DEVICE, libusb.LIBUSB_ERROR_NOT_FOUND):
libusb.libusb_close(self.handle)
self.handle = self.reopen_device(_desc.idVendor, _desc.idProduct, bus_number)
if checked(libusb.libusb_kernel_driver_active)(self.handle, 0):
checked(libusb.libusb_detach_kernel_driver)(self.handle, 0)
elif reset_rc < 0:
raise RuntimeError(f"libusb_reset_device: {ctypes.string_at(libusb.libusb_strerror(reset_rc)).decode()}")
checked(libusb.libusb_reset_device)(self.handle)
# Set configuration and claim interface
checked(libusb.libusb_set_configuration)(self.handle, 1)
checked(libusb.libusb_claim_interface)(self.handle, 0)
checked(libusb.libusb_set_interface_alt_setting)(self.handle, 0, 0)
if use_bot:
checked(libusb.libusb_set_interface_alt_setting)(self.handle, 0, 0)
self._tag = 0
else:
checked(libusb.libusb_set_interface_alt_setting)(self.handle, 0, 1)
def control_write(self, request:int, value:int=0, index:int=0, data:bytes=b'', timeout:int=1000):
assert len(data) <= len(self._ctrl_mv)
self._ctrl_mv[:len(data)] = data
assert checked(libusb.libusb_control_transfer)(self.handle, 0x40, request, value, index, self._ctrl_buf, len(data), timeout) == len(data)
# Clear any stalled endpoints
all_eps = (self.ep_data_out, self.ep_data_in, self.ep_stat_in, self.ep_cmd_out)
for ep in all_eps: checked(libusb.libusb_clear_halt)(self.handle, ep)
def control_read(self, request:int, length:int, value:int=0, index:int=0, timeout:int=1000) -> memoryview:
assert length <= len(self._ctrl_mv)
assert checked(libusb.libusb_control_transfer)(self.handle, 0xC0, request, value, index, self._ctrl_buf, length, timeout) == length
return self._ctrl_mv[:length]
# Allocate streams
stream_eps = (ctypes.c_uint8 * 3)(self.ep_data_out, self.ep_data_in, self.ep_stat_in)
checked(libusb.libusb_alloc_streams)(self.handle, self.max_streams * len(stream_eps), stream_eps, len(stream_eps))
def bulk_write(self, payload:bytes, timeout:int=1000):
if len(payload) > len(self._bulk_mv): self._bulk_buf, self._bulk_mv = alloc_cbuffer(len(payload))
self._bulk_mv[:len(payload)] = payload
checked(libusb.libusb_bulk_transfer, "bulk OUT 0x02 failed") \
(self.handle, 0x02, self._bulk_buf, len(payload), self._transferred, timeout)
assert self._transferred.value == len(payload), f"bulk OUT short write: {self._transferred.value}/{len(payload)} bytes"
# Base cmd
cmd_template = bytes([0x01, 0x00, 0x00, 0x01, *([0] * 12), 0xE4, 0x24, 0x00, 0xB2, 0x1A, 0x00, 0x00, 0x00, *([0] * 8)])
def bulk_read(self, length:int, timeout:int=1000) -> memoryview:
if length > len(self._bulk_mv): self._bulk_buf, self._bulk_mv = alloc_cbuffer(length)
checked(libusb.libusb_bulk_transfer, "bulk IN 0x81 failed")(self.handle, 0x81, self._bulk_buf, length, self._transferred, timeout)
return self._bulk_mv[:self._transferred.value]
# Init pools
self.tr = {ep: [libusb.libusb_alloc_transfer(0) for _ in range(self.max_streams)] for ep in all_eps}
self.buf_cmd = [(ctypes.c_uint8 * len(cmd_template))(*cmd_template) for _ in range(self.max_streams)]
self.buf_stat = [(ctypes.c_uint8 * 64)() for _ in range(self.max_streams)]
self.buf_data_in = [(ctypes.c_uint8 * 0x1000)() for _ in range(self.max_streams)]
self.buf_data_out = [(ctypes.c_uint8 * 0x80000)() for _ in range(self.max_streams)]
self.buf_data_out_mvs = [to_mv(ctypes.addressof(self.buf_data_out[i]), 0x80000) for i in range(self.max_streams)]
for slot in range(self.max_streams): struct.pack_into(">B", self.buf_cmd[slot], 3, slot + 1)
def _prep_transfer(self, tr, ep, stream_id, buf, length):
tr.contents.dev_handle, tr.contents.endpoint, tr.contents.length, tr.contents.buffer = self.handle, ep, length, buf
tr.contents.status, tr.contents.flags, tr.contents.timeout, tr.contents.num_iso_packets = 0xff, 0, 1000, 0
tr.contents.type = (libusb.LIBUSB_TRANSFER_TYPE_BULK_STREAM if stream_id is not None else libusb.LIBUSB_TRANSFER_TYPE_BULK)
if stream_id is not None: libusb.libusb_transfer_set_stream_id(tr, stream_id)
return tr
def _submit_and_wait(self, cmds):
for tr in cmds: checked(libusb.libusb_submit_transfer)(tr)
running = len(cmds)
while running:
checked(libusb.libusb_handle_events)(USB3.ctx())
running = len(cmds)
for tr in cmds:
if tr.contents.status == libusb.LIBUSB_TRANSFER_COMPLETED: running -= 1
elif tr.contents.status != 0xFF: raise RuntimeError(f"EP 0x{tr.contents.endpoint:02X} error: {tr.contents.status}")
def _bulk_out(self, ep: int, payload: bytes, timeout: int = 1000):
if len(payload) > len(self._bulk_out_mv): self._bulk_out_buf, self._bulk_out_mv = alloc_cbuffer(len(payload))
self._bulk_out_mv[:len(payload)] = payload
checked(libusb.libusb_bulk_transfer, f"bulk OUT 0x{ep:02X} failed")(self.handle, ep, self._bulk_out_buf, len(payload), self._transferred, timeout)
assert self._transferred.value == len(payload), f"bulk OUT short write on 0x{ep:02X}: {self._transferred.value}/{len(payload)} bytes"
def _bulk_in(self, ep: int, length: int, timeout: int = 1000) -> memoryview:
if length > len(self._bulk_in_mv): self._bulk_in_buf, self._bulk_in_mv = alloc_cbuffer(length)
checked(libusb.libusb_bulk_transfer, f"bulk IN 0x{ep:02X} failed")(self.handle, ep, self._bulk_in_buf, length, self._transferred, timeout)
return self._bulk_in_mv[:self._transferred.value]
def send_batch(self, cdbs:list[bytes], idata:list[int]|None=None, odata:list[bytes|None]|None=None) -> list[bytes|None]:
idata, odata = idata or [0] * len(cdbs), odata or [None] * len(cdbs)
results:list[bytes|None] = []
tr_window, op_window = [], []
for idx, (cdb, rlen, send_data) in enumerate(zip(cdbs, idata, odata)):
if self.use_bot:
dir_in = rlen > 0
data_len = rlen if dir_in else (len(send_data) if send_data is not None else 0)
assert not (rlen > 0 and send_data is not None), "BOT mode only supports either read or write per command"
# CBW
self._tag += 1
flags = 0x80 if dir_in else 0x00
cbw = struct.pack("<IIIBBB", 0x43425355, self._tag, data_len, flags, 0, len(cdb)) + cdb + b"\x00" * (16 - len(cdb))
self._bulk_out(self.ep_data_out, cbw)
# DAT
if dir_in:
results.append(bytes(self._bulk_in(self.ep_data_in, rlen)))
else:
if send_data is not None:
self._bulk_out(self.ep_data_out, send_data)
results.append(None)
# CSW
sig, rtag, residue, status = struct.unpack("<IIIB", self._bulk_in(self.ep_data_in, 13, timeout=2000))
assert sig == 0x53425355, f"Bad CSW signature 0x{sig:08X}, expected 0x53425355"
assert rtag == self._tag, f"CSW tag mismatch: got {rtag}, expected {self._tag}"
assert status == 0, f"SCSI command failed, CSW status=0x{status:02X}, residue={residue}"
else:
# allocate slot and stream. stream is 1-based
slot, stream = idx % self.max_streams, (idx % self.max_streams) + 1
# build cmd packet
self.buf_cmd[slot][16:16+len(cdb)] = list(cdb)
# cmd + stat transfers
tr_window.append(self._prep_transfer(self.tr[self.ep_cmd_out][slot], self.ep_cmd_out, None, self.buf_cmd[slot], len(self.buf_cmd[slot])))
tr_window.append(self._prep_transfer(self.tr[self.ep_stat_in][slot], self.ep_stat_in, stream, self.buf_stat[slot], 64))
if rlen:
if rlen > len(self.buf_data_in[slot]): self.buf_data_in[slot] = (ctypes.c_uint8 * round_up(rlen, 0x1000))()
tr_window.append(self._prep_transfer(self.tr[self.ep_data_in][slot], self.ep_data_in, stream, self.buf_data_in[slot], rlen))
if send_data is not None:
if len(send_data) > len(self.buf_data_out[slot]):
self.buf_data_out[slot] = (ctypes.c_uint8 * len(send_data))()
self.buf_data_out_mvs[slot] = to_mv(ctypes.addressof(self.buf_data_out[slot]), len(send_data))
self.buf_data_out_mvs[slot][:len(send_data)] = bytes(send_data)
tr_window.append(self._prep_transfer(self.tr[self.ep_data_out][slot], self.ep_data_out, stream, self.buf_data_out[slot], len(send_data)))
op_window.append((idx, slot, rlen))
if (idx + 1 == len(cdbs)) or len(op_window) >= self.max_streams:
self._submit_and_wait(tr_window)
for idx, slot, rlen in op_window: results.append(bytes(self.buf_data_in[slot][:rlen]) if rlen else None)
tr_window = []
return results
@dataclasses.dataclass(frozen=True)
class WriteOp: addr:int; data:bytes; ignore_cache:bool=True # noqa: E702
@dataclasses.dataclass(frozen=True)
class ReadOp: addr:int; size:int # noqa: E702
@dataclasses.dataclass(frozen=True)
class ScsiWriteOp: data:bytes; lba:int=0 # noqa: E702
# NOTE: keep it for flash.py
def send_batch(self, cdbs:list[bytes], odata:list[bytes|None]|None=None):
for cdb, data in zip(cdbs, odata or [None] * len(cdbs)):
self.bulk_write(struct.pack("<IIIBBB16s", 0x43425355, tag:=next(self._tags), len(data) if data is not None else 0, 0, 0, len(cdb), cdb))
if data is not None: self.bulk_write(data)
sig, rtag, _, status = struct.unpack("<IIIB", self.bulk_read(13, timeout=2000))
assert (sig, rtag, status) == (0x53425355, tag, 0)
class CustomASM24Controller:
def __init__(self, usb:USB3|None=None):
if not usb:
devs = USB3.list_devices(0xADD1, 0x0001) + USB3.list_devices(0x3801, 0x0001)
assert len(devs), "no ASM24 controller found"
self.usb = USB3(devs[0][0], 0x81, 0x83, 0x02, 0x04, use_bot=True)
else: self.usb = usb
self._pci_cacheable: list[tuple[int, int]] = []
self._pci_cache: dict[int, int|None] = {}
def __init__(self, usb:USB3):
self.usb = usb
self._f0_out_buf, self._f0_out_mv = alloc_cbuffer(0x1000) # for f0 and e4, allocate big enough for e4
self._f0_in_buf, _ = alloc_cbuffer(8)
if (ltssm := self.read(0xB450, 1)[0]) != 0x78:
self.set_pcie_power(True)
grace_period = time.monotonic() + 5.0
while time.monotonic() < grace_period and (ltssm := self.read(0xB450, 1)[0]) != 0x78:
time.sleep(0.1)
# Custom firmware now boots with PCIe off. Power it on before probing the link.
ltssm = self.read(0xB450, 1)[0]
if ltssm != 0x78: self.set_pcie_power(True)
ltssm = self.read(0xB450, 1)[0]
if ltssm != 0x78: raise RuntimeError(f"PCIe link not up (LTSSM=0x{ltssm:02X}), custom firmware not ready")
def set_pcie_power(self, enabled:bool, timeout:int=10000):
checked(libusb.libusb_control_transfer,
f"F3 PCIe power {'on' if enabled else 'off'} failed")(self.usb.handle, 0x40, 0xF3, int(enabled), 0, None, 0, timeout)
# === PCIe TLP via 0xF0 vendor command ===
def set_pcie_power(self, enabled:bool, timeout:int=10000): self.usb.control_write(0xF3, value=int(enabled), timeout=timeout)
def _f0_out(self, fmt_type:int, byte_en:int, address:int, value:int, mode:int=0):
struct.pack_into('<III', self._f0_out_mv, 0, address & 0xFFFFFFFF, address >> 32, value)
ret = libusb.libusb_control_transfer(self.usb.handle, 0x40, 0xF0, fmt_type | (byte_en << 8), mode & 0x03, self._f0_out_buf, 12, 5000)
assert ret == 12, f"F0 OUT failed: {ret}"
self.usb.control_write(0xF0, fmt_type | (byte_en << 8), mode & 0x03, struct.pack('<III', address & 0xFFFFFFFF, address >> 32, value), 5000)
def _f0_in(self) -> tuple[int, int, int]:
ret = libusb.libusb_control_transfer(self.usb.handle, 0xC0, 0xF0, 0, 0, self._f0_in_buf, 8, 5000)
assert ret == 8, f"F0 IN failed: {ret}"
return struct.unpack_from('<I', self._f0_in_buf, 0)[0], (self._f0_in_buf[4] >> 5) & 0x7, self._f0_in_buf[7]
def _is_pci_cacheable(self, addr:int) -> bool: return any(x <= addr <= x + sz for x, sz in self._pci_cacheable)
data = self.usb.control_read(0xF0, 8, timeout=5000)
return struct.unpack_from('<I', data)[0], (data[4] >> 5) & 0x7, data[7]
def pcie_request(self, fmt_type:int, address:int, value:int|None=None, size:int=4, cnt:int=10):
if fmt_type == 0x60 and size == 4 and self._is_pci_cacheable(address) and self._pci_cache.get(address) == value: return
assert size > 0 and size <= 4, f"Invalid size {size}"
if DEBUG >= 5: print("pcie_request", hex(fmt_type), hex(address), value, size)
offset = address & 0x3
byte_en = ((1 << size) - 1) << offset
self._pci_cache[address] = value if size == 4 and fmt_type == 0x60 else None
self._f0_out(fmt_type, byte_en, address & ~0x3, (value << (8 * offset)) if value is not None else 0)
# Fast path: memory writes and messages don't return completions (same logic as ASM24Controller).
# Fast path: memory writes and messages don't return completions.
if ((fmt_type & 0b11011111) == 0b01000000) or ((fmt_type & 0b10111000) == 0b00110000): return
# Read TLPs and config writes: read completion via 0xF0 IN. Retry on error/timeout.
data, cpl_status, ret_status = self._f0_in()
if ret_status != 0:
time.sleep(0.001) # TODO: this sleep is very picky
if cnt > 0:
return self.pcie_request(fmt_type, address, value, size, cnt=cnt-1)
if cnt > 0: return self.pcie_request(fmt_type, address, value, size, cnt=cnt-1)
raise RuntimeError(f"TLP error after retries: ret_status={ret_status}, address={address:#x}")
if cpl_status:
@@ -281,219 +135,69 @@ class CustomASM24Controller:
address = (bus << 24) | (dev << 19) | (fn << 16) | (byte_addr & 0xfff)
return self.pcie_request(fmt_type, address, value, size)
def pcie_mem_req(self, address:int, value:int|None=None, size:int=4):
return self.pcie_request(0x60 if value is not None else 0x20, address, value, size)
def pcie_mem_write(self, address:int, values:list[int], size:int):
def pcie_mem_write(self, address:int, data:bytes):
"""Streaming PCIe memory write via 0xF0 mode 1 + bulk OUT. Data is little-endian dwords on the wire."""
if not values: return
self._f0_out(0x60, 0x0F, address, len(values), mode=1)
self.usb._bulk_out(0x02, struct.pack(f'<{len(values)}I', *values))
if not data: return
assert len(data) % 4 == 0, f"pcie_mem_write requires 4-byte aligned size, got {len(data)}"
self._f0_out(0x60, 0x0F, address, len(data) // 4, mode=1)
self.usb.bulk_write(data)
def pcie_mem_read(self, address:int, nbytes:int) -> bytes:
def pcie_mem_read(self, address:int, nbytes:int) -> memoryview:
"""Streaming PCIe memory read via 0xF0 mode 2 + bulk IN. Returns little-endian bytes."""
assert nbytes % 4 == 0, f"pcie_mem_read requires 4-byte aligned size, got {nbytes}"
self._f0_out(0x20, 0x0F, address, nbytes // 4, mode=2)
return self.usb._bulk_in(0x81, nbytes, timeout=30000)
return self.usb.bulk_read(nbytes, timeout=30000)
# === XDATA read/write (0xE4/0xE5 vendor control transfers) ===
def read(self, base_addr:int, length:int, **kwargs) -> bytes:
def read(self, base_addr:int, length:int) -> bytes:
"""Read from chip XDATA via vendor control IN (bRequest=0xE4). wValue=addr, wLength=size."""
result = b''
for off in range(0, length, 0xFF):
chunk = min(0xFF, length - off)
ret = libusb.libusb_control_transfer(self.usb.handle, 0xC0, 0xE4, base_addr + off, 0, self._f0_out_buf, chunk, 1000)
assert ret == chunk, f"read(0x{base_addr + off:04X}, {chunk}) failed: {ret}"
result += bytes(self._f0_out_buf[:ret])
return result[:length]
result += self.usb.control_read(0xE4, chunk, value=base_addr + off)
return result
def write(self, base_addr:int, data:bytes, **kwargs):
def write(self, base_addr:int, data:bytes):
"""Write to chip XDATA via vendor control OUT (bRequest=0xE5). wValue=addr, wIndex=val."""
for off, val in enumerate(data):
checked(libusb.libusb_control_transfer,
f"write(0x{base_addr + off:04X}, 0x{val:02X}) failed")(self.usb.handle, 0x40, 0xE5, base_addr + off, val, None, 0, 1000)
for off, val in enumerate(data): self.usb.control_write(0xE5, value=base_addr + off, index=val)
def scsi_write(self, buf:bytes, lba:int=0):
def scsi_write(self, buf:bytes):
"""Write to SRAM via 0xF2 vendor command + bulk OUT."""
buf_padded = buf + b'\x00' * (round_up(len(buf), 512) - len(buf))
sectors = len(buf_padded) // 512
num_slots = round_up(len(buf_padded), 0x4000) // 0x4000 # 16KB per slot
# 0xF2 OUT: wValue=sectors, wIndex=start_slot|(num_slots<<8)
num_slots = ceildiv(len(buf_padded), 0x4000) # 16KB per slot
windex = (num_slots & 0xFF) << 8
checked(libusb.libusb_control_transfer, "F2 setup failed")(self.usb.handle, 0x40, 0xF2, sectors, windex, None, 0, 1000)
self.usb._bulk_out(0x02, buf_padded)
self.usb.control_write(0xF2, value=sectors, index=windex)
self.usb.bulk_write(buf_padded)
def scsi_read_arm(self, size:int):
windex = (ceildiv(size, 0x4000) & 0xFF) << 8
checked(libusb.libusb_control_transfer,
"F2 read arm failed")(self.usb.handle, 0x40, 0xF2, (ceildiv(size, 512) & 0x7FFF) | 0x8000, windex, None, 0, 1000)
self.usb.control_write(0xF2, value=(ceildiv(size, 512) & 0x7FFF) | 0x8000, index=windex)
def scsi_read(self, size:int) -> memoryview: return self.usb._bulk_in(0x81, round_up(size, 512), timeout=10000)[:size]
class ASM24Controller:
def __init__(self, usb:USB3|None=None):
if not usb:
devs = USB3.list_devices(0xADD1, 0x0001) + USB3.list_devices(0x3801, 0x0001)
assert len(devs), "no ASM24 controller found"
self.usb = USB3(devs[0][0], 0x81, 0x83, 0x02, 0x04, use_bot=bool(getenv("USE_BOT", 0)))
else: self.usb = usb
self._cache: dict[int, int|None] = {}
self._pci_cacheable: list[tuple[int, int]] = []
self._pci_cache: dict[int, int|None] = {}
# Init controller.
self.exec_ops([WriteOp(0x54b, b' '), WriteOp(0x54e, b'\x04'), WriteOp(0x5a8, b'\x02'), WriteOp(0x5f8, b'\x04'),
WriteOp(0x7ec, b'\x01\x00\x00\x00'), WriteOp(0xc422, b'\x02'), WriteOp(0x0, b'\x33')])
def exec_ops(self, ops:Sequence[WriteOp|ReadOp|ScsiWriteOp]):
cdbs:list[bytes] = []
idata:list[int] = []
odata:list[bytes|None] = []
def _add_req(cdb:bytes, i:int, o:bytes|None):
nonlocal cdbs, idata, odata
cdbs, idata, odata = cdbs + [cdb], idata + [i], odata + [o]
for op in ops:
if isinstance(op, WriteOp):
for off, value in enumerate(op.data):
addr = ((op.addr + off) & 0x1FFFF) | 0x500000
if not op.ignore_cache and self._cache.get(addr) == value: continue
_add_req(struct.pack('>BBBHB', 0xE5, value, addr >> 16, addr & 0xFFFF, 0), 0, None)
self._cache[addr] = value
elif isinstance(op, ReadOp):
assert op.size <= 0xff
addr = (op.addr & 0x1FFFF) | 0x500000
_add_req(struct.pack('>BBBHB', 0xE4, op.size, addr >> 16, addr & 0xFFFF, 0), op.size, None)
for i in range(op.size): self._cache[addr + i] = None
elif isinstance(op, ScsiWriteOp):
sectors = round_up(len(op.data), 512) // 512
_add_req(struct.pack('>BBQIBB', 0x8A, 0, op.lba, sectors, 0, 0), 0, op.data+b'\x00'*((sectors*512)-len(op.data)))
return self.usb.send_batch(cdbs, idata, odata)
def write(self, base_addr:int, data:bytes, ignore_cache:bool=True): return self.exec_ops([WriteOp(base_addr, data, ignore_cache)])
def scsi_write(self, buf:bytes, lba:int=0):
if len(buf) > 0x4000: buf += b'\x00' * (round_up(len(buf), 0x10000) - len(buf))
for i in range(0, len(buf), 0x10000):
self.exec_ops([ScsiWriteOp(buf[i:i+0x10000], lba), WriteOp(0x171, b'\xff\xff\xff', ignore_cache=True)])
self.exec_ops([WriteOp(0xce6e, b'\x00\x00', ignore_cache=True)])
if len(buf) > 0x4000:
for i in range(4): self.exec_ops([WriteOp(0xce40 + i, b'\x00', ignore_cache=True)])
def read(self, base_addr:int, length:int, stride:int=0xff) -> bytes:
parts = self.exec_ops([ReadOp(base_addr + off, min(stride, length - off)) for off in range(0, length, stride)])
return b''.join(p or b'' for p in parts)[:length]
def _is_pci_cacheable(self, addr:int) -> bool: return any(x <= addr <= x + sz for x, sz in self._pci_cacheable)
def pcie_prep_request(self, fmt_type:int, address:int, value:int|None=None, size:int=4) -> list[WriteOp]:
if fmt_type == 0x60 and size == 4 and self._is_pci_cacheable(address) and self._pci_cache.get(address) == value: return []
assert fmt_type >> 8 == 0 and size > 0 and size <= 4, f"Invalid fmt_type {fmt_type} or size {size}"
if DEBUG >= 5: print("pcie_request", hex(fmt_type), hex(address), value, size)
masked_address, offset = address & 0xFFFFFFFC, address & 0x3
assert size + offset <= 4 and (value is None or value >> (8 * size) == 0)
self._pci_cache[address] = value if size == 4 and fmt_type == 0x60 else None
return ([WriteOp(0xB220, struct.pack('>I', value << (8 * offset)), ignore_cache=False)] if value is not None else []) + \
[WriteOp(0xB218, struct.pack('>I', masked_address), ignore_cache=False), WriteOp(0xB21c, struct.pack('>I', address>>32), ignore_cache=False),
WriteOp(0xB217, bytes([((1 << size) - 1) << offset]), ignore_cache=False), WriteOp(0xB210, bytes([fmt_type]), ignore_cache=False),
WriteOp(0xB254, b"\x0f", ignore_cache=True), WriteOp(0xB296, b"\x04", ignore_cache=True)]
def pcie_request(self, fmt_type, address, value=None, size=4, cnt=10):
self.exec_ops(self.pcie_prep_request(fmt_type, address, value, size))
# Fast path for write requests
if ((fmt_type & 0b11011111) == 0b01000000) or ((fmt_type & 0b10111000) == 0b00110000): return
while (stat:=self.read(0xB296, 1)[0]) & 2 == 0:
if stat & 1:
self.write(0xB296, bytes([0x01]))
if cnt > 0: return self.pcie_request(fmt_type, address, value, size, cnt=cnt-1)
assert stat == 2, f"stat read 2 was {stat}"
# Retrieve completion data from Link Status (0xB22A, 0xB22B)
b284 = self.read(0xB284, 1)[0]
completion = struct.unpack('>H', self.read(0xB22A, 2))
# Validate completion status based on PCIe request typ
# Completion TLPs for configuration requests always have a byte count of 4.
assert completion[0] & 0xfff == (4 if (fmt_type & 0xbe == 0x04) else size)
# Extract completion status field
status = (completion[0] >> 13) & 0x7
# Handle completion errors or inconsistencies
if status or ((fmt_type & 0xbe == 0x04) and (((value is None) and (not (b284 & 0x01))) or ((value is not None) and (b284 & 0x01)))):
status_map = {0b001: f"Unsupported Request: invalid address/function (target might not be reachable): {address:#x}",
0b100: "Completer Abort: abort due to internal error", 0b010: "Configuration Request Retry Status: configuration space busy"}
raise RuntimeError(f"TLP status: {status_map.get(status, 'Reserved (0b{:03b})'.format(status))}")
if value is None: return (struct.unpack('>I', self.read(0xB220, 4))[0] >> (8 * (address & 0x3))) & ((1 << (8 * size)) - 1)
def pcie_cfg_req(self, byte_addr, bus=1, dev=0, fn=0, value=None, size=4):
assert byte_addr >> 12 == 0 and bus >> 8 == 0 and dev >> 5 == 0 and fn >> 3 == 0, f"Invalid byte_addr {byte_addr}, bus {bus}, dev {dev}, fn {fn}"
fmt_type = (0x44 if value is not None else 0x4) | int(bus > 0)
address = (bus << 24) | (dev << 19) | (fn << 16) | (byte_addr & 0xfff)
return self.pcie_request(fmt_type, address, value, size)
def pcie_mem_req(self, address, value=None, size=4): return self.pcie_request(0x60 if value is not None else 0x20, address, value, size)
def pcie_mem_write(self, address, values, size):
ops = [self.pcie_prep_request(0x60, address + i * size, value, size) for i, value in enumerate(values)]
# Send in batches of 4 for OSX and 16 for Linux (benchmarked values)
for i in range(0, len(ops), bs:=(4 if OSX else 16)): self.exec_ops(list(itertools.chain.from_iterable(ops[i:i+bs])))
def scsi_read(self, size:int) -> memoryview: return self.usb.bulk_read(round_up(size, 512), timeout=10000)[:size]
class USBMMIOInterface(MMIOInterface):
def __init__(self, usb, addr, size, fmt, pcimem=True): # pylint: disable=super-init-not-called
self.usb, self.addr, self.nbytes, self.fmt, self.pcimem, self.el_sz = usb, addr, size, fmt, pcimem, struct.calcsize(fmt)
self.usb, self.addr, self.nbytes, self.fmt, self.el_sz, self.pcimem = usb, addr, size, fmt, struct.calcsize(fmt), pcimem
def __getitem__(self, index): return self._access_items(index)
def __setitem__(self, index, val): self._access_items(index, val)
def _off_from_index(self, index):
if isinstance(index, slice): return ((index.start or 0) * self.el_sz, ((index.stop or len(self))-(index.start or 0)) * self.el_sz)
return (index * self.el_sz, self.el_sz)
def _access_items(self, index, val=None):
if isinstance(index, slice): return self._acc((index.start or 0) * self.el_sz, ((index.stop or len(self))-(index.start or 0)) * self.el_sz, val)
return self._acc_one(index * self.el_sz, self.el_sz, val) if self.pcimem else self._acc(index * self.el_sz, self.el_sz, val)
def __getitem__(self, index):
off, sz = self._off_from_index(index)
if self.pcimem:
assert sz % 4 == 0 and off % 4 == 0, f"pcie_mem_read requires 4-byte aligned access, got off={off}, sz={sz}"
data = self.usb.pcie_mem_read(self.addr + off, sz)
else: data = self.usb.scsi_read(sz) if self.addr == 0xf000 else self.usb.read(self.addr + off, sz)
return int.from_bytes(data, "little") if sz == self.el_sz else data
def __setitem__(self, index, data):
off, _ = self._off_from_index(index)
data = struct.pack(self.fmt, data) if isinstance(data, int) else bytes(data)
if not self.pcimem: self.usb.scsi_write(data) if self.addr == 0xf000 else self.usb.write(self.addr + off, data)
else: self.usb.pcie_mem_write(self.addr+off, data)
def view(self, offset:int=0, size:int|None=None, fmt=None):
return USBMMIOInterface(self.usb, self.addr+offset, size or (self.nbytes - offset), fmt=fmt or self.fmt, pcimem=self.pcimem)
def _acc_size(self, sz): return next(x for x in [('I', 4), ('H', 2), ('B', 1)] if sz % x[1] == 0)
def _acc_one(self, off, sz, val=None):
upper = 0 if sz < 8 else self.usb.pcie_mem_req(self.addr + off + 4, val if val is None else (val >> 32), 4)
lower = self.usb.pcie_mem_req(self.addr + off, val if val is None else val & 0xffffffff, min(sz, 4))
if val is None: return lower | (upper << 32)
def _acc(self, off, sz, data=None):
if data is None: # read op
if not self.pcimem:
if self.addr == 0xf000 and hasattr(self.usb, 'scsi_read'): return self.usb.scsi_read(sz)
return int.from_bytes(self.usb.read(self.addr + off, sz), "little") if sz == self.el_sz else self.usb.read(self.addr + off, sz)
# Fast path: streaming PCIe read if controller supports it
if hasattr(self.usb, 'pcie_mem_read') and sz >= 4 and sz % 4 == 0:
return self.usb.pcie_mem_read(self.addr + off, sz)
acc, acc_size = self._acc_size(sz)
return bytes(array.array(acc, [self._acc_one(off + i * acc_size, acc_size) for i in range(sz // acc_size)]))
# write op
data = struct.pack(self.fmt, data) if isinstance(data, int) else bytes(data)
if not self.pcimem:
# Fast path for writing into buffer 0xf000
use_cache = 0xa800 <= self.addr <= 0xb000
return self.usb.scsi_write(bytes(data)) if self.addr == 0xf000 else self.usb.write(self.addr + off, bytes(data), ignore_cache=not use_cache)
_, acc_sz = self._acc_size(len(data) * struct.calcsize(self.fmt))
self.usb.pcie_mem_write(self.addr+off, [int.from_bytes(data[i:i+acc_sz], "little") for i in range(0, len(data), acc_sz)], acc_sz)
return USBMMIOInterface(self.usb, self.addr+offset, self.nbytes-offset if size is None else size, fmt=fmt or self.fmt, pcimem=self.pcimem)
if DEV.interface.startswith("MOCK"): from test.mockgpu.usb import MockUSB3 as USB3 # type: ignore # noqa: F811