"""Unit tests for the RM power-limit interface (fake RM, no hardware). Standalone (no pytest required): python tests/test_rm_power.py Also works under pytest if available. Ports the test battery from LACT PR #1205 (ilya-zlobintsev/LACT): layout discovery, NVML cross-validation, write minimality, readback verification, and failure restoration. """ import os import sys sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..")) from nvcurve.hal.rm_power import ( # noqa: E402 _CTRL_GPU_GET_ATTACHED_IDS, _CTRL_GPU_GET_ID_INFO_V2, _CTRL_GPU_GET_PCI_INFO, _PWR_GET_CONTROL, _PWR_GET_INFO, _PWR_SET_CONTROL, EXTENDED_LAYOUT, LEGACY_LAYOUT, PciLocation, PowerLimitBounds, RmPowerError, _u32, probe, resolve_gpu_instance, set_limit, ) PASS = 0 FAIL = 0 def check(name: str, cond: bool) -> None: global PASS, FAIL if cond: PASS += 1 print(f" PASS {name}") else: FAIL += 1 print(f" FAIL {name}") BOUNDS = PowerLimitBounds(min_mw=250_000, default_mw=300_000, max_mw=325_000) class FakeRm: """In-memory fake of the RM power-limit client (both wire layouts).""" def __init__(self, layout, current: int) -> None: self.layout = layout self.control = bytearray(layout.control_size) self.control[0:8] = bytes([0xFF, 0, 0, 0, 1, 0, 0, 0]) self.control[layout.request_at - 4 : layout.request_at] = bytes( [0x67, 0x67, 0, 0] ) self.control[layout.request_at : layout.request_at + 4] = current.to_bytes( 4, "little" ) self.control[layout.client_at] = 0xFE self.reads: list[tuple[int, int]] = [] self.writes: list[bytes] = [] self.fail_first_write = False self.fail_readback = False self.fail_restore = False def query(self, cmd: int, data: bytearray) -> None: if cmd == _PWR_GET_INFO: self.reads.append((cmd, len(data))) if len(data) != self.layout.info_size: raise RmPowerError("Unsupported INFO size") data[0:8] = bytes([0xFF, 0, 0, 0, 1, 0, 0, 0]) for index, value in enumerate([250_000, 300_000, 325_000]): offset = self.layout.info_min_at + 4 * index data[offset : offset + 4] = value.to_bytes(4, "little") elif cmd == _PWR_GET_CONTROL: self.reads.append((cmd, len(data))) if len(data) != self.layout.control_size: raise RmPowerError("Unsupported CONTROL size") if data[self.layout.client_at] != 0xFE: raise AssertionError("unexpected client selector on GET") if self.fail_readback and len(self.writes) == 1: raise RmPowerError("readback unavailable") data[:] = self.control elif cmd == _PWR_SET_CONTROL: if len(data) != self.layout.control_size: raise AssertionError("bad SET size") if data[4:8] != (1).to_bytes(4, "little"): raise AssertionError("bad SET mask") if data[self.layout.client_at] != 0xFE: raise AssertionError("bad SET client selector") # Only the request field may differ from the current state. for i, (a, b) in enumerate(zip(data, self.control, strict=True)): if self.layout.request_at <= i < self.layout.request_at + 4: continue if a != b: raise AssertionError(f"SET modified byte {i:#x}") self.writes.append(bytes(data)) if self.fail_restore and len(self.writes) > 1: raise RmPowerError("restore unavailable") self.control[:] = data if self.fail_first_write and len(self.writes) == 1: raise RmPowerError("SET failed after modifying hardware") else: raise AssertionError(f"unexpected command {cmd:#x}") def test_detects_both_layouts_with_gets_without_a_driver_version() -> None: for layout in (EXTENDED_LAYOUT, LEGACY_LAYOUT): rm = FakeRm(layout, 250_000) support = probe(BOUNDS, 250_000, rm.query) check( f"{layout.name}: detected", support.bounds == BOUNDS and support.layout == layout, ) expected = ( [(_PWR_GET_INFO, 0x924), (_PWR_GET_CONTROL, 0x328)] if layout == EXTENDED_LAYOUT else [ (_PWR_GET_INFO, 0x924), (_PWR_GET_INFO, 0x488), (_PWR_GET_CONTROL, 0x188), ] ) check(f"{layout.name}: GETs only, expected sequence", rm.reads == expected) check(f"{layout.name}: no writes during discovery", rm.writes == []) def test_unknown_layout_and_nvml_mismatches_never_write() -> None: calls: list[tuple[int, int]] = [] def failing(cmd: int, data: bytearray) -> None: calls.append((cmd, len(data))) raise RmPowerError("Unsupported payload") try: probe(BOUNDS, 250_000, failing) check("unknown layout rejected", False) except RmPowerError: check("unknown layout rejected", True) check( "unknown layout: only GETs attempted", calls == [(_PWR_GET_INFO, 0x924), (_PWR_GET_INFO, 0x488)], ) for layout in (EXTENDED_LAYOUT, LEGACY_LAYOUT): rm = FakeRm(layout, 250_000) try: probe(BOUNDS, 300_000, rm.query) check(f"{layout.name}: current mismatch rejected", False) except RmPowerError: check(f"{layout.name}: current mismatch rejected", True) other_bounds = PowerLimitBounds( min_mw=BOUNDS.min_mw, default_mw=BOUNDS.default_mw, max_mw=350_000 ) try: probe(other_bounds, 250_000, rm.query) check(f"{layout.name}: bounds mismatch rejected", False) except RmPowerError: check(f"{layout.name}: bounds mismatch rejected", True) check(f"{layout.name}: no writes on mismatch", rm.writes == []) def test_rejects_unrecognized_headers_masks_and_client_values() -> None: for layout in (EXTENDED_LAYOUT, LEGACY_LAYOUT): for at, value in [(0, 0), (4, 3), (layout.client_at, 0xF8)]: rm = FakeRm(layout, 250_000) rm.control[at] = value try: probe(BOUNDS, 250_000, rm.query) check(f"{layout.name}: bad header/client rejected", False) except RmPowerError: check(f"{layout.name}: bad header/client rejected", True) check(f"{layout.name}: no writes on bad header", rm.writes == []) for current in (0, 0xFFFFFFFF): rm = FakeRm(layout, current) try: probe(BOUNDS, current, rm.query) check(f"{layout.name}: empty request rejected", False) except RmPowerError: check(f"{layout.name}: empty request rejected", True) # The extended layout has additional mask words. Accepting only its low # word would allow an unexpected client to be included in a later SET. rm = FakeRm(EXTENDED_LAYOUT, 250_000) rm.control[8] = 1 try: probe(BOUNDS, 250_000, rm.query) check("extended: nonzero mask word rejected", False) except RmPowerError: check("extended: nonzero mask word rejected", True) check("extended: no writes on mask violation", rm.writes == []) def test_changes_only_fe_request_and_keeps_vbios_maximum() -> None: for layout in (EXTENDED_LAYOUT, LEGACY_LAYOUT): rm = FakeRm(layout, 250_000) support = probe(BOUNDS, 250_000, rm.query) for cap in (150_000, 30_000, 250_000): set_limit(cap, support, rm.query) check( f"{layout.name}: set {cap} mW", _u32(rm.control, layout.request_at) == cap, ) writes = len(rm.writes) for cap in (0, 29_999, 325_001, 350_000, 0xFFFFFFFF): try: set_limit(cap, support, rm.query) check(f"{layout.name}: out-of-range {cap} rejected", False) except RmPowerError: check(f"{layout.name}: out-of-range {cap} rejected", True) check( f"{layout.name}: no writes for out-of-range caps", len(rm.writes) == writes, ) def test_restores_previous_below_minimum_request_after_set_or_readback_failure() -> ( None ): for layout in (EXTENDED_LAYOUT, LEGACY_LAYOUT): for fail_set in (False, True): rm = FakeRm(layout, 100_000) original = bytes(rm.control) rm.fail_first_write = fail_set rm.fail_readback = not fail_set support = probe(BOUNDS, 100_000, rm.query) try: set_limit(150_000, support, rm.query) check(f"{layout.name}: failure reported", False) except RmPowerError: check(f"{layout.name}: failure reported", True) check(f"{layout.name}: restore issued", len(rm.writes) == 2) check( f"{layout.name}: previous request restored", bytes(rm.control) == original, ) def test_reports_restore_failure_and_rejects_wrong_client_before_writing() -> None: for layout in (EXTENDED_LAYOUT, LEGACY_LAYOUT): rm = FakeRm(layout, 100_000) rm.fail_first_write = True rm.fail_restore = True support = probe(BOUNDS, 100_000, rm.query) try: set_limit(150_000, support, rm.query) check(f"{layout.name}: restore failure reported", False) except RmPowerError as exc: check( f"{layout.name}: restore failure reported", "restoration also failed" in str(exc), ) rm = FakeRm(layout, 250_000) rm.control[layout.client_at] = 0xF8 try: set_limit(150_000, support, rm.query) check(f"{layout.name}: wrong client rejected", False) except RmPowerError: check(f"{layout.name}: wrong client rejected", True) check(f"{layout.name}: no writes for wrong client", rm.writes == []) # ── PCI identity → RM instance resolution ──────────────────────────────────── def test_resolves_pci_identity_when_minor_and_rm_orders_differ() -> None: # This host has Ada at minor 5/RM 4 and the 5090 at minor 4/RM 5. # IDs are opaque and enumeration order must not select the device. pci = PciLocation(domain=0, bus=0x0D, dev=0, func=0) instances = resolve_gpu_instance(pci, lambda cmd, data: _fake_root(cmd, data)) check("resolves by PCI identity", instances == (5, 2)) def _fake_root(cmd: int, data: bytearray) -> None: if cmd == _CTRL_GPU_GET_ATTACHED_IDS: data[0:4] = (0x2E00).to_bytes(4, "little") data[4:8] = (0x0D00).to_bytes(4, "little") elif cmd == _CTRL_GPU_GET_PCI_INFO: gpu_id = _u32(data, 0) bus = 0x2E if gpu_id == 0x2E00 else 0x0D data[8:10] = bus.to_bytes(2, "little") elif cmd == _CTRL_GPU_GET_ID_INFO_V2: if _u32(data, 0) != 0x0D00: raise AssertionError("unexpected gpu id in ID_INFO_V2") data[8:12] = (5).to_bytes(4, "little") data[12:16] = (2).to_bytes(4, "little") else: raise AssertionError(f"unexpected command {cmd:#x}") def test_does_not_fall_back_to_another_gpu_when_pci_is_missing() -> None: pci = PciLocation(domain=1, bus=0x0D, dev=0, func=0) def query(cmd: int, data: bytearray) -> None: if cmd == _CTRL_GPU_GET_ATTACHED_IDS: data[0:4] = (0x0D00).to_bytes(4, "little") elif cmd == _CTRL_GPU_GET_PCI_INFO: data[8:10] = (0x0D).to_bytes(2, "little") else: raise AssertionError("must not allocate a GPU from another PCI domain") try: resolve_gpu_instance(pci, query) check("foreign PCI domain rejected", False) except RmPowerError: check("foreign PCI domain rejected", True) def test_rejects_nonzero_pci_function() -> None: pci = PciLocation(domain=0, bus=0x0D, dev=0, func=1) try: resolve_gpu_instance(pci, lambda cmd, data: None) check("nonzero function rejected", False) except RmPowerError: check("nonzero function rejected", True) def test_propagates_rm_query_failure() -> None: pci = PciLocation(domain=0, bus=0x0D, dev=0, func=0) try: resolve_gpu_instance( pci, lambda cmd, data: (_ for _ in ()).throw(RmPowerError("RM unavailable")) ) check("RM query failure propagated", False) except RmPowerError as exc: check("RM query failure propagated", "RM unavailable" in str(exc)) def main() -> int: tests = [ test_detects_both_layouts_with_gets_without_a_driver_version, test_unknown_layout_and_nvml_mismatches_never_write, test_rejects_unrecognized_headers_masks_and_client_values, test_changes_only_fe_request_and_keeps_vbios_maximum, test_restores_previous_below_minimum_request_after_set_or_readback_failure, test_reports_restore_failure_and_rejects_wrong_client_before_writing, test_resolves_pci_identity_when_minor_and_rm_orders_differ, test_does_not_fall_back_to_another_gpu_when_pci_is_missing, test_rejects_nonzero_pci_function, test_propagates_rm_query_failure, ] for t in tests: print(f"== {t.__name__} ==") t() print(f"\n{PASS} passed, {FAIL} failed") return 1 if FAIL else 0 if __name__ == "__main__": sys.exit(main())