From 588d87b40a3630a2d77bb6346777ce67fe07890e Mon Sep 17 00:00:00 2001 From: Lucas Manning Date: Tue, 16 Jan 2024 17:25:57 -0800 Subject: [PATCH] Add a unit test for sentry sysfs PCI mirroring. PiperOrigin-RevId: 599006437 --- pkg/sentry/fsimpl/sys/pci.go | 8 +-- pkg/sentry/fsimpl/sys/sys.go | 12 +++-- pkg/sentry/fsimpl/sys/sys_test.go | 84 +++++++++++++++++++++++++++++-- 3 files changed, 92 insertions(+), 12 deletions(-) diff --git a/pkg/sentry/fsimpl/sys/pci.go b/pkg/sentry/fsimpl/sys/pci.go index 513bfcba2..e6c50b6e2 100644 --- a/pkg/sentry/fsimpl/sys/pci.go +++ b/pkg/sentry/fsimpl/sys/pci.go @@ -53,10 +53,12 @@ var ( } ) -// Creates TPU devices' symlinks under /sys/class/. TPU deivce type that are not present on host willl be ignored. +// Creates TPU devices' symlinks under /sys/class/. TPU device types that are +// not present on host will be ignored. +// // TPU v4 symlinks are created at /sys/class/accel/accel#. // TPU v5 symlinks go to /sys/class/vfio-dev/vfio#. -func (fs *filesystem) newDeviceClassDir(ctx context.Context, creds *auth.Credentials, tpuDeviceTypes []string) (map[string]map[string]kernfs.Inode, error) { +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{} pciDents, err := hostDirEntries(pciMainBusDevicePath) if err != nil { @@ -87,7 +89,7 @@ func (fs *filesystem) newDeviceClassDir(ctx context.Context, creds *auth.Credent } // Create /sys/bus/pci/devices symlinks. -func (fs *filesystem) newPCIDevicesDir(ctx context.Context, creds *auth.Credentials) (map[string]kernfs.Inode, error) { +func (fs *filesystem) newBusPCIDevicesDir(ctx context.Context, creds *auth.Credentials, pciMainBusDevicePath string) (map[string]kernfs.Inode, error) { pciDevicesDir := map[string]kernfs.Inode{} pciDents, err := hostDirEntries(pciMainBusDevicePath) if err != nil { diff --git a/pkg/sentry/fsimpl/sys/sys.go b/pkg/sentry/fsimpl/sys/sys.go index 11daf2046..c01329c12 100644 --- a/pkg/sentry/fsimpl/sys/sys.go +++ b/pkg/sentry/fsimpl/sys/sys.go @@ -58,6 +58,9 @@ type InternalData struct { // EnableTPUProxyPaths is whether to populate sysfs paths used by hardware // accelerators. EnableTPUProxyPaths bool + // PCIDevicePathPrefix is a prefix for the PCI device paths. It is useful for + // unit testing. + PCIDevicePathPrefix string } // filesystem implements vfs.FilesystemImpl. @@ -130,17 +133,18 @@ func (fsType FilesystemType) GetFilesystem(ctx context.Context, vfsObj *vfs.Virt idata := opts.InternalData.(*InternalData) productName = idata.ProductName if idata.EnableTPUProxyPaths { - deviceToIommuGroup, err := pciDeviceIOMMUGroups(iommuGroupSysPath) + deviceToIommuGroup, err := pciDeviceIOMMUGroups(path.Join(idata.PCIDevicePathPrefix, iommuGroupSysPath)) if err != nil { return nil, nil, err } - pciMainBusSub, err := fs.mirrorPCIBusDeviceDir(ctx, creds, pciMainBusDevicePath, deviceToIommuGroup) + pciPath := path.Join(idata.PCIDevicePathPrefix, pciMainBusDevicePath) + pciMainBusSub, err := fs.mirrorPCIBusDeviceDir(ctx, creds, pciPath, deviceToIommuGroup) if err != nil { return nil, nil, err } devicesSub["pci0000:00"] = fs.newDir(ctx, creds, defaultSysDirMode, pciMainBusSub) - deviceDirs, err := fs.newDeviceClassDir(ctx, creds, []string{accelDevice, vfioDevice}) + deviceDirs, err := fs.newDeviceClassDir(ctx, creds, []string{accelDevice, vfioDevice}, pciPath) if err != nil { return nil, nil, err } @@ -148,7 +152,7 @@ func (fsType FilesystemType) GetFilesystem(ctx context.Context, vfsObj *vfs.Virt for tpuDeviceType, symlinkDir := range deviceDirs { classSub[tpuDeviceType] = fs.newDir(ctx, creds, defaultSysDirMode, symlinkDir) } - pciDevicesSub, err := fs.newPCIDevicesDir(ctx, creds) + pciDevicesSub, err := fs.newBusPCIDevicesDir(ctx, creds, pciPath) if err != nil { return nil, nil, err } diff --git a/pkg/sentry/fsimpl/sys/sys_test.go b/pkg/sentry/fsimpl/sys/sys_test.go index d70374b68..06f4fee5d 100644 --- a/pkg/sentry/fsimpl/sys/sys_test.go +++ b/pkg/sentry/fsimpl/sys/sys_test.go @@ -16,6 +16,8 @@ package sys_test import ( "fmt" + "os" + "path" "testing" "github.com/google/go-cmp/cmp" @@ -27,7 +29,7 @@ import ( "gvisor.dev/gvisor/pkg/sentry/vfs" ) -func newTestSystem(t *testing.T) *testutil.System { +func newTestSystem(t *testing.T, pciTestDir string) *testutil.System { k, err := testutil.Boot() if err != nil { t.Fatalf("Failed to create test kernel: %v", err) @@ -38,7 +40,16 @@ func newTestSystem(t *testing.T) *testutil.System { AllowUserMount: true, }) - mns, err := k.VFS().NewMountNamespace(ctx, creds, "", sys.Name, &vfs.MountOptions{}, nil) + mountOpts := &vfs.MountOptions{ + GetFilesystemOptions: vfs.GetFilesystemOptions{ + InternalData: &sys.InternalData{ + EnableTPUProxyPaths: pciTestDir != "", + PCIDevicePathPrefix: pciTestDir, + }, + }, + } + + mns, err := k.VFS().NewMountNamespace(ctx, creds, "", sys.Name, mountOpts, nil) if err != nil { t.Fatalf("Failed to create new mount namespace: %v", err) } @@ -46,7 +57,7 @@ func newTestSystem(t *testing.T) *testutil.System { } func TestReadCPUFile(t *testing.T) { - s := newTestSystem(t) + s := newTestSystem(t, "" /*pciTestDir*/) defer s.Destroy() k := kernel.KernelFromContext(s.Ctx) maxCPUCores := k.ApplicationCores() @@ -71,7 +82,7 @@ func TestReadCPUFile(t *testing.T) { } func TestSysRootContainsExpectedEntries(t *testing.T) { - s := newTestSystem(t) + s := newTestSystem(t, "" /*pciTestDir*/) defer s.Destroy() pop := s.PathOpAtRoot("/") s.AssertAllDirentTypes(s.ListDirents(pop), map[string]testutil.DirentType{ @@ -90,7 +101,7 @@ func TestSysRootContainsExpectedEntries(t *testing.T) { func TestCgroupMountpointExists(t *testing.T) { // Note: The mountpoint is only created if cgroups are available. - s := newTestSystem(t) + s := newTestSystem(t, "" /*pciTestDir*/) defer s.Destroy() pop := s.PathOpAtRoot("/fs") s.AssertAllDirentTypes(s.ListDirents(pop), map[string]testutil.DirentType{ @@ -99,3 +110,66 @@ func TestCgroupMountpointExists(t *testing.T) { pop = s.PathOpAtRoot("/fs/cgroup") s.AssertAllDirentTypes(s.ListDirents(pop), map[string]testutil.DirentType{ /*empty*/ }) } + +// Check that sysfs creates the required PCI paths for V4 TPUs. +func TestEnableTPUProxyPathsV4(t *testing.T) { + // Set up the fs tree that will be mirrored in the sentry. + pciTestDir := t.TempDir() + accelPath := path.Join(pciTestDir, "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(pciTestDir, "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(pciTestDir, "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) + } + + s := newTestSystem(t, pciTestDir) + defer s.Destroy() + + pop := s.PathOpAtRoot("/devices/pci0000:00") + s.AssertAllDirentTypes(s.ListDirents(pop), map[string]testutil.DirentType{ + "0000:00:04.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{ + "0000:00:04.0": linux.DT_LNK, + "device": linux.DT_LNK, + "chip_model": linux.DT_REG, + "device_owner": linux.DT_REG, + "pci_address": linux.DT_REG, + }) + pop = s.PathOpAtRoot("/bus/pci/devices") + s.AssertAllDirentTypes(s.ListDirents(pop), map[string]testutil.DirentType{ + "0000:00:04.0": linux.DT_LNK, + }) + pop = s.PathOpAtRoot("/class/accel") + s.AssertAllDirentTypes(s.ListDirents(pop), map[string]testutil.DirentType{ + "accel0": linux.DT_LNK, + }) +}