Fix pci device mirroring so it doesn't overwrite device directories.

PiperOrigin-RevId: 650320903
This commit is contained in:
Lucas Manning
2024-07-08 11:39:54 -07:00
committed by gVisor bot
parent 9d1849029e
commit 67596b46a7
2 changed files with 33 additions and 25 deletions
+3 -1
View File
@@ -60,6 +60,9 @@ var (
// TPU v5 symlinks go to /sys/class/vfio-dev/vfio#.
func (fs *filesystem) newDeviceClassDir(ctx context.Context, creds *auth.Credentials, tpuDeviceTypes []string, pciMainBusDevicePath string) (map[string]map[string]kernfs.Inode, error) {
dirs := map[string]map[string]kernfs.Inode{}
for _, tpuDeviceType := range tpuDeviceTypes {
dirs[tpuDeviceType] = map[string]kernfs.Inode{}
}
pciDents, err := hostDirEntries(pciMainBusDevicePath)
if err != nil {
return nil, err
@@ -67,7 +70,6 @@ func (fs *filesystem) newDeviceClassDir(ctx context.Context, creds *auth.Credent
for _, pciDent := range pciDents {
for _, tpuDeviceType := range tpuDeviceTypes {
subPath := path.Join(pciMainBusDevicePath, pciDent, tpuDeviceType)
dirs[tpuDeviceType] = map[string]kernfs.Inode{}
deviceDents, err := hostDirEntries(subPath)
if err != nil {
// Skips the path that doesn't exist.
+30 -24
View File
@@ -119,38 +119,41 @@ func TestCgroupMountpointExists(t *testing.T) {
func TestEnableTPUProxyPathsV4(t *testing.T) {
// Set up the fs tree that will be mirrored in the sentry.
sysfsTestDir := t.TempDir()
accelPath := path.Join(sysfsTestDir, "sys", "devices", "pci0000:00", "0000:00:04.0", "accel", "accel0")
if err := os.MkdirAll(accelPath, 0755); err != nil {
t.Fatalf("Failed to create accel directory: %v", err)
}
if err := os.Symlink(path.Join("..", "..", "..", "0000:00:04.0"), path.Join(accelPath, "0000:00:04.0")); err != nil {
t.Fatalf("Failed to symlink accel directory: %v", err)
}
if err := os.Symlink(path.Join("..", "..", "..", "0000:00:04.0"), path.Join(accelPath, "device")); err != nil {
t.Fatalf("Failed to symlink accel device directory: %v", err)
}
if _, err := os.Create(path.Join(accelPath, "chip_model")); err != nil {
t.Fatalf("Failed to create chip_model: %v", err)
}
if _, err := os.Create(path.Join(accelPath, "device_owner")); err != nil {
t.Fatalf("Failed to create device_owner: %v", err)
}
if _, err := os.Create(path.Join(accelPath, "pci_address")); err != nil {
t.Fatalf("Failed to create pci_address: %v", err)
}
busPath := path.Join(sysfsTestDir, "sys", "bus", "pci", "devices")
if err := os.MkdirAll(busPath, 0755); err != nil {
t.Fatalf("Failed to create bus directory: %v", err)
}
if err := os.Symlink(path.Join("..", "..", "..", "devices", "pci0000:00", "0000:00:04.0"), path.Join(busPath, "0000:00:04.0")); err != nil {
t.Fatalf("Failed to symlink bus directory: %v", err)
}
classAccelPath := path.Join(sysfsTestDir, "sys", "class", "accel")
if err := os.MkdirAll(classAccelPath, 0755); err != nil {
t.Fatalf("Failed to create accel directory: %v", err)
}
if err := os.Symlink(path.Join("..", "..", "devices", "pci0000:00", "0000:00:04.0", "accel", "accel0"), path.Join(classAccelPath, "accel0")); err != nil {
t.Fatalf("Failed to symlink accel directory: %v", err)
for i, pciAddress := range []string{"0000:00:04.0", "0000:00:05.0"} {
accelDev := fmt.Sprintf("accel%d", i)
accelPath := path.Join(sysfsTestDir, "sys", "devices", "pci0000:00", pciAddress, "accel", accelDev)
if err := os.MkdirAll(accelPath, 0755); err != nil {
t.Fatalf("Failed to create accel directory: %v", err)
}
if err := os.Symlink(path.Join("..", "..", "..", pciAddress), path.Join(accelPath, pciAddress)); err != nil {
t.Fatalf("Failed to symlink accel directory: %v", err)
}
if err := os.Symlink(path.Join("..", "..", "..", pciAddress), path.Join(accelPath, "device")); err != nil {
t.Fatalf("Failed to symlink accel device directory: %v", err)
}
if _, err := os.Create(path.Join(accelPath, "chip_model")); err != nil {
t.Fatalf("Failed to create chip_model: %v", err)
}
if _, err := os.Create(path.Join(accelPath, "device_owner")); err != nil {
t.Fatalf("Failed to create device_owner: %v", err)
}
if _, err := os.Create(path.Join(accelPath, "pci_address")); err != nil {
t.Fatalf("Failed to create pci_address: %v", err)
}
if err := os.Symlink(path.Join("..", "..", "..", "devices", "pci0000:00", pciAddress), path.Join(busPath, pciAddress)); err != nil {
t.Fatalf("Failed to symlink bus directory: %v", err)
}
if err := os.Symlink(path.Join("..", "..", "devices", "pci0000:00", pciAddress, "accel", accelDev), path.Join(classAccelPath, accelDev)); err != nil {
t.Fatalf("Failed to symlink accel directory: %v", err)
}
}
s := newTestSystem(t, sysfsTestDir)
@@ -159,6 +162,7 @@ func TestEnableTPUProxyPathsV4(t *testing.T) {
pop := s.PathOpAtRoot("/devices/pci0000:00")
s.AssertAllDirentTypes(s.ListDirents(pop), map[string]testutil.DirentType{
"0000:00:04.0": linux.DT_DIR,
"0000:00:05.0": linux.DT_DIR,
})
pop = s.PathOpAtRoot("/devices/pci0000:00/0000:00:04.0/accel/accel0")
s.AssertAllDirentTypes(s.ListDirents(pop), map[string]testutil.DirentType{
@@ -171,10 +175,12 @@ func TestEnableTPUProxyPathsV4(t *testing.T) {
pop = s.PathOpAtRoot("/bus/pci/devices")
s.AssertAllDirentTypes(s.ListDirents(pop), map[string]testutil.DirentType{
"0000:00:04.0": linux.DT_LNK,
"0000:00:05.0": linux.DT_LNK,
})
pop = s.PathOpAtRoot("/class/accel")
s.AssertAllDirentTypes(s.ListDirents(pop), map[string]testutil.DirentType{
"accel0": linux.DT_LNK,
"accel1": linux.DT_LNK,
})
}