Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
25 changes: 19 additions & 6 deletions dau_utils/pci_runtime_pm.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,11 @@

import argparse
import subprocess
import sys
from collections.abc import Sequence
from dataclasses import dataclass
from pathlib import Path

DEFAULT_DEVICE_PATTERNS = ("Thunderbolt", "JHL", "10ee:7011", "Xilinx")
DEFAULT_SYSFS_ROOT = Path("/sys/bus/pci/devices")


Expand All @@ -16,7 +16,7 @@ class RuntimePmWrite:
value: str


def discover_pci_devices(lspci_output: str, *, patterns: Sequence[str] = DEFAULT_DEVICE_PATTERNS) -> tuple[str, ...]:
def discover_pci_devices(lspci_output: str, *, patterns: Sequence[str] = ()) -> tuple[str, ...]:
matches: list[str] = []
for line in lspci_output.splitlines():
if any(pattern in line for pattern in patterns):
Expand All @@ -43,7 +43,7 @@ def plan_runtime_pm_writes(mode: str, devices: Sequence[str], *, sysfs_root: Pat
def apply_runtime_pm_writes(writes: Sequence[RuntimePmWrite]) -> tuple[RuntimePmWrite, ...]:
applied: list[RuntimePmWrite] = []
for write in writes:
if write.path.exists() and write.path.parent.exists():
if write.path.exists():
write.path.write_text(f"{write.value}\n")
applied.append(write)
return tuple(applied)
Expand All @@ -59,17 +59,30 @@ def main(argv: Sequence[str] | None = None) -> int:
parser.add_argument("--dry-run", action="store_true", help="Print writes without applying them")
args = parser.parse_args(argv)

patterns = tuple(args.pattern) if args.pattern else DEFAULT_DEVICE_PATTERNS
devices = tuple(args.device) or discover_pci_devices(args.lspci_output if args.lspci_output is not None else _lspci_output(), patterns=patterns)
patterns = tuple(args.pattern)
if args.device:
devices: tuple[str, ...] = tuple(args.device)
elif patterns:
lspci_output = args.lspci_output if args.lspci_output is not None else _lspci_output()
devices = discover_pci_devices(lspci_output, patterns=patterns)
else:
devices = ()
writes = plan_runtime_pm_writes(args.mode, devices, sysfs_root=args.sysfs_root)

if args.dry_run:
for write in writes:
print(f"write {write.path} {write.value}")
return 0

for write in apply_runtime_pm_writes(writes):
applied = apply_runtime_pm_writes(writes)
for write in applied:
print(f"wrote {write.path} {write.value}")
skipped = tuple(write for write in writes if write not in applied)
for write in skipped:
print(f"skipped {write.path} (missing)", file=sys.stderr)
if skipped:
print(f"applied {len(applied)} of {len(writes)} runtime PM writes", file=sys.stderr)
return 1
return 0


Expand Down
80 changes: 77 additions & 3 deletions dau_utils/tests/test_pci_runtime_pm.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,15 +14,20 @@
"""


def test_discovers_thunderbolt_and_xilinx_devices_from_lspci_output() -> None:
assert discover_pci_devices(LSPCI_OUTPUT) == (
def test_discovers_devices_matching_explicit_patterns() -> None:
assert discover_pci_devices(LSPCI_OUTPUT, patterns=("Thunderbolt", "JHL", "10ee:7011")) == (
"0000:00:07.0",
"0000:00:0d.2",
"0000:02:00.0",
"0000:04:00.0",
)


def test_no_patterns_discovers_no_devices() -> None:
assert discover_pci_devices(LSPCI_OUTPUT) == ()
assert discover_pci_devices(LSPCI_OUTPUT, patterns=()) == ()


def test_runtime_pm_write_plan_maps_hold_and_release_to_sysfs_knobs() -> None:
root = Path("/sys/bus/pci/devices")

Expand Down Expand Up @@ -50,14 +55,83 @@ def test_cli_dry_run_prints_hold_writes_for_explicit_device(capsys) -> None:


def test_cli_dry_run_can_discover_devices_from_lspci_fixture(capsys) -> None:
exit_code = main(["release", "--dry-run", "--lspci-output", LSPCI_OUTPUT])
exit_code = main(
[
"release",
"--dry-run",
"--pattern",
"Thunderbolt",
"--pattern",
"JHL",
"--pattern",
"10ee:7011",
"--lspci-output",
LSPCI_OUTPUT,
]
)

assert exit_code == 0
lines = capsys.readouterr().out.splitlines()
assert lines[0] == "write /sys/bus/pci/devices/0000:00:07.0/power/control auto"
assert lines[-1] == "write /sys/bus/pci/devices/0000:04:00.0/d3cold_allowed 1"


def test_cli_without_patterns_matches_no_devices(capsys) -> None:
exit_code = main(["hold"])

assert exit_code == 0
assert capsys.readouterr().out == ""


def test_cli_missing_sysfs_paths_surface_skips_and_fail(tmp_path, capsys) -> None:
exit_code = main(["hold", "--device", "0000:04:00.0", "--sysfs-root", str(tmp_path)])

assert exit_code == 1
captured = capsys.readouterr()
assert captured.out == ""
err_lines = captured.err.splitlines()
assert err_lines == [
f"skipped {tmp_path}/0000:04:00.0/power/control (missing)",
f"skipped {tmp_path}/0000:04:00.0/d3cold_allowed (missing)",
"applied 0 of 2 runtime PM writes",
]


def test_cli_partial_apply_reports_skip_and_fails(tmp_path, capsys) -> None:
device_root = tmp_path / "0000:04:00.0"
(device_root / "power").mkdir(parents=True)
control = device_root / "power" / "control"
control.write_text("auto\n")
# d3cold_allowed intentionally absent

exit_code = main(["hold", "--device", "0000:04:00.0", "--sysfs-root", str(tmp_path)])

assert exit_code == 1
assert control.read_text() == "on\n"
captured = capsys.readouterr()
assert captured.out.splitlines() == [f"wrote {control} on"]
assert captured.err.splitlines() == [
f"skipped {device_root}/d3cold_allowed (missing)",
"applied 1 of 2 runtime PM writes",
]


def test_cli_applies_present_sysfs_paths(tmp_path, capsys) -> None:
device_root = tmp_path / "0000:04:00.0"
(device_root / "power").mkdir(parents=True)
control = device_root / "power" / "control"
d3cold = device_root / "d3cold_allowed"
control.write_text("auto\n")
d3cold.write_text("1\n")

exit_code = main(["hold", "--device", "0000:04:00.0", "--sysfs-root", str(tmp_path)])

assert exit_code == 0
assert control.read_text() == "on\n"
assert d3cold.read_text() == "0\n"
assert capsys.readouterr().err == ""


def test_module_entrypoint_runs_cli_for_uninstalled_checkout(capsys, monkeypatch) -> None:
monkeypatch.setattr(sys, "argv", ["pci_runtime_pm", "hold", "--device", "0000:04:00.0", "--dry-run"])
monkeypatch.delitem(sys.modules, "dau_utils.pci_runtime_pm", raising=False)
Expand Down
Loading