From 44d987d9849110511492260d583d5bf8ad96d04b Mon Sep 17 00:00:00 2001 From: Lucas Manning Date: Fri, 27 Oct 2023 15:17:02 -0700 Subject: [PATCH] Add sanitization checks for TPU ioctl parameters and rename TPU device methods. The methods are renamed so that it's more clear they are explicitly for TPUV4, not any arbitrary accel device. PiperOrigin-RevId: 577317824 --- pkg/abi/tpu/BUILD | 11 ++ pkg/abi/tpu/tpu.go | 137 ++++++++++++++++++ pkg/sentry/devices/accel/BUILD | 7 +- pkg/sentry/devices/accel/device.go | 27 +++- pkg/sentry/devices/accel/gasket.go | 69 ++++++++- .../{accel_unsafe.go => gasket_unsafe.go} | 0 pkg/sentry/devices/accel/seccomp_filters.go | 28 ---- .../devices/accel/{accel.go => tpu_v4.go} | 37 +++-- .../accel/{accel_mmap.go => tpu_v4_mmap.go} | 14 +- runsc/boot/BUILD | 1 + runsc/boot/vfs.go | 29 +++- runsc/cmd/util/BUILD | 1 + runsc/cmd/util/tpu.go | 4 +- 13 files changed, 302 insertions(+), 63 deletions(-) create mode 100644 pkg/abi/tpu/BUILD create mode 100644 pkg/abi/tpu/tpu.go rename pkg/sentry/devices/accel/{accel_unsafe.go => gasket_unsafe.go} (100%) rename pkg/sentry/devices/accel/{accel.go => tpu_v4.go} (85%) rename pkg/sentry/devices/accel/{accel_mmap.go => tpu_v4_mmap.go} (85%) diff --git a/pkg/abi/tpu/BUILD b/pkg/abi/tpu/BUILD new file mode 100644 index 000000000..cc158e303 --- /dev/null +++ b/pkg/abi/tpu/BUILD @@ -0,0 +1,11 @@ +load("//tools:defs.bzl", "go_library") + +package(default_applicable_licenses = ["//:license"]) + +licenses(["notice"]) + +go_library( + name = "tpu", + srcs = ["tpu.go"], + visibility = ["//pkg/sentry:internal"], +) diff --git a/pkg/abi/tpu/tpu.go b/pkg/abi/tpu/tpu.go new file mode 100644 index 000000000..4ca7e4352 --- /dev/null +++ b/pkg/abi/tpu/tpu.go @@ -0,0 +1,137 @@ +// Copyright 2023 The gVisor Authors. +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Package tpu defines constants used to interact with TPUs. The constants are +// derived from those listed in https://github.com/tensorflow/tpu/blob/master/tools/driver/drivers/char/tpu_common +package tpu + +const ( + // SizeOfTPUV4InterruptList is the total number of valid + // (BAR Index, Register Offset) pairs. + SizeOfTPUV4InterruptList = uint64(45) + + // NumberOfTPUV4PageTables is the number of gasket page tables used by the + // TPU V4 driver. + NumberOfTPUV4PageTables = uint64(1) + + // TPUV4DeviceID is the PCI device ID of TPU V4 hardware. + TPUV4DeviceID = 0x005E + + // SizeOfTPUV4liteInterruptList is the total number of valid + // (BAR Index, Register Offset) pairs. + SizeOfTPUV4liteInterruptList = uint64(37) + + // NumberOfTPUV4litePageTables is the number of gasket page tables used by the + // TPU V4 driver + NumberOfTPUV4litePageTables = uint64(1) + + // TPUV4liteDeviceID is the PCI device ID of TPU V4lite hardware. + TPUV4liteDeviceID = 0x0056 +) + +// TPUV4InterruptsMap maps BAR indices to valid register offsets. +var ( + TPUV4InterruptsMap = map[uint64]map[uint64]struct{}{ + 2: map[uint64]struct{}{ + 0x15b0008: struct{}{}, + 0x15b0000: struct{}{}, + 0x16b0008: struct{}{}, + 0x16b0000: struct{}{}, + 0x17b0008: struct{}{}, + 0x17b0000: struct{}{}, + 0x18b0008: struct{}{}, + 0x18b0000: struct{}{}, + 0x19b0020: struct{}{}, + 0x19b0000: struct{}{}, + 0x19b0008: struct{}{}, + 0x19b0010: struct{}{}, + 0x19b0018: struct{}{}, + 0x1ab0020: struct{}{}, + 0x1ab0000: struct{}{}, + 0x1ab0008: struct{}{}, + 0x1ab0010: struct{}{}, + 0x1ab0018: struct{}{}, + 0x4720000: struct{}{}, + 0x1bb0000: struct{}{}, + 0x1bb0008: struct{}{}, + 0x1bb0010: struct{}{}, + 0x1bb0018: struct{}{}, + 0x90000: struct{}{}, + 0xb0000: struct{}{}, + 0xd0000: struct{}{}, + 0xf0000: struct{}{}, + 0x110000: struct{}{}, + 0x130000: struct{}{}, + 0x150000: struct{}{}, + 0x170000: struct{}{}, + 0x190000: struct{}{}, + 0x1b0000: struct{}{}, + 0x1d0000: struct{}{}, + 0x1f0000: struct{}{}, + 0x210000: struct{}{}, + 0x230000: struct{}{}, + 0x250000: struct{}{}, + 0x270000: struct{}{}, + 0x290000: struct{}{}, + 0x2b0000: struct{}{}, + 0x2d0000: struct{}{}, + 0x2f0000: struct{}{}, + 0x310000: struct{}{}, + 0x4720018: struct{}{}, + }, + } + + // TPUV4liteInterruptsMap maps BAR indices to valid register offsets. + TPUV4liteInterruptsMap = map[uint64]map[uint64]struct{}{ + 2: map[uint64]struct{}{ + 0x19b0020: struct{}{}, + 0x19b0000: struct{}{}, + 0x19b0008: struct{}{}, + 0x19b0010: struct{}{}, + 0x19b0018: struct{}{}, + 0x1ab0020: struct{}{}, + 0x1ab0000: struct{}{}, + 0x1ab0008: struct{}{}, + 0x1ab0010: struct{}{}, + 0x1ab0018: struct{}{}, + 0x4720000: struct{}{}, + 0x1bb0000: struct{}{}, + 0x1bb0008: struct{}{}, + 0x1bb0010: struct{}{}, + 0x1bb0018: struct{}{}, + 0x90000: struct{}{}, + 0xb0000: struct{}{}, + 0xd0000: struct{}{}, + 0xf0000: struct{}{}, + 0x110000: struct{}{}, + 0x130000: struct{}{}, + 0x150000: struct{}{}, + 0x170000: struct{}{}, + 0x190000: struct{}{}, + 0x1b0000: struct{}{}, + 0x1d0000: struct{}{}, + 0x1f0000: struct{}{}, + 0x210000: struct{}{}, + 0x230000: struct{}{}, + 0x250000: struct{}{}, + 0x270000: struct{}{}, + 0x290000: struct{}{}, + 0x2b0000: struct{}{}, + 0x2d0000: struct{}{}, + 0x2f0000: struct{}{}, + 0x310000: struct{}{}, + 0x4720018: struct{}{}, + }, + } +) diff --git a/pkg/sentry/devices/accel/BUILD b/pkg/sentry/devices/accel/BUILD index 29a9c2388..5f9d97e63 100644 --- a/pkg/sentry/devices/accel/BUILD +++ b/pkg/sentry/devices/accel/BUILD @@ -6,19 +6,20 @@ licenses(["notice"]) go_library( name = "accel", srcs = [ - "accel.go", - "accel_mmap.go", - "accel_unsafe.go", "devaddr_range.go", "devaddr_set.go", "device.go", "gasket.go", + "gasket_unsafe.go", "seccomp_filters.go", + "tpu_v4.go", + "tpu_v4_mmap.go", ], visibility = ["//pkg/sentry:internal"], deps = [ "//pkg/abi/gasket", "//pkg/abi/linux", + "//pkg/abi/tpu", "//pkg/cleanup", "//pkg/context", "//pkg/errors/linuxerr", diff --git a/pkg/sentry/devices/accel/device.go b/pkg/sentry/devices/accel/device.go index 56339c853..ec09dbfdf 100644 --- a/pkg/sentry/devices/accel/device.go +++ b/pkg/sentry/devices/accel/device.go @@ -20,26 +20,31 @@ import ( "golang.org/x/sys/unix" "gvisor.dev/gvisor/pkg/abi/linux" "gvisor.dev/gvisor/pkg/context" + "gvisor.dev/gvisor/pkg/errors/linuxerr" "gvisor.dev/gvisor/pkg/fdnotifier" "gvisor.dev/gvisor/pkg/sentry/fsimpl/devtmpfs" + "gvisor.dev/gvisor/pkg/sentry/kernel" "gvisor.dev/gvisor/pkg/sentry/vfs" "gvisor.dev/gvisor/pkg/sync" ) -// accelDevice implements vfs.Device for /dev/accel[0-9]+. +// tpuV4Device implements vfs.Device for /dev/accel[0-9]+. // // +stateify savable -type accelDevice struct { +type tpuV4Device struct { mu sync.Mutex minor uint32 + lite bool // +checklocks:mu openWriteFDs uint32 // +checklocks:mu devAddrSet DevAddrSet + // +checklocks:mu + owner *kernel.ThreadGroup } -func (dev *accelDevice) Open(ctx context.Context, mnt *vfs.Mount, vfsd *vfs.Dentry, opts vfs.OpenOptions) (*vfs.FileDescription, error) { +func (dev *tpuV4Device) Open(ctx context.Context, mnt *vfs.Mount, vfsd *vfs.Dentry, opts vfs.OpenOptions) (*vfs.FileDescription, error) { dev.mu.Lock() defer dev.mu.Unlock() hostPath := fmt.Sprintf("/dev/accel%d", dev.minor) @@ -48,7 +53,7 @@ func (dev *accelDevice) Open(ctx context.Context, mnt *vfs.Mount, vfsd *vfs.Dent ctx.Warningf("accelDevice: failed to open host %s: %v", hostPath, err) return nil, err } - fd := &accelFD{ + fd := &tpuV4FD{ hostFD: int32(hostFD), device: dev, } @@ -66,6 +71,13 @@ func (dev *accelDevice) Open(ctx context.Context, mnt *vfs.Mount, vfsd *vfs.Dent if vfs.MayWriteFileWithOpenFlags(opts.Flags) { dev.openWriteFDs++ } + if dev.owner == nil { + t := kernel.TaskFromContext(ctx) + if t == nil { + return nil, linuxerr.ESRCH + } + dev.owner = t.ThreadGroup() + } return &fd.vfsfd, nil } @@ -74,9 +86,10 @@ func CreateDevtmpfsFile(ctx context.Context, dev *devtmpfs.Accessor, num uint32) return dev.CreateDeviceFile(ctx, fmt.Sprintf("accel%d", num), vfs.CharDevice, linux.ACCEL_MAJOR, num, 0666) } -// Register registers all devices implemented by this package in vfsObj. -func Register(vfsObj *vfs.VirtualFilesystem, minor uint32) error { - return vfsObj.RegisterDevice(vfs.CharDevice, linux.ACCEL_MAJOR, minor, &accelDevice{ +// RegisterTPUV4Device registers all devices implemented by this package in vfsObj. +func RegisterTPUV4Device(vfsObj *vfs.VirtualFilesystem, minor uint32, lite bool) error { + return vfsObj.RegisterDevice(vfs.CharDevice, linux.ACCEL_MAJOR, minor, &tpuV4Device{ + lite: lite, minor: minor, }, &vfs.RegisterDeviceOptions{ GroupName: "accel", diff --git a/pkg/sentry/devices/accel/gasket.go b/pkg/sentry/devices/accel/gasket.go index 98e59f1bc..a916739be 100644 --- a/pkg/sentry/devices/accel/gasket.go +++ b/pkg/sentry/devices/accel/gasket.go @@ -20,6 +20,7 @@ import ( "golang.org/x/sys/unix" "gvisor.dev/gvisor/pkg/abi/gasket" "gvisor.dev/gvisor/pkg/abi/linux" + "gvisor.dev/gvisor/pkg/abi/tpu" "gvisor.dev/gvisor/pkg/cleanup" "gvisor.dev/gvisor/pkg/context" "gvisor.dev/gvisor/pkg/errors/linuxerr" @@ -30,12 +31,20 @@ import ( "gvisor.dev/gvisor/pkg/sentry/mm" ) -func gasketMapBufferIoctl(ctx context.Context, t *kernel.Task, hostFd int32, fd *accelFD, paramsAddr hostarch.Addr) (uintptr, error) { +func gasketMapBufferIoctl(ctx context.Context, t *kernel.Task, hostFd int32, fd *tpuV4FD, paramsAddr hostarch.Addr) (uintptr, error) { var userIoctlParams gasket.GasketPageTableIoctl if _, err := userIoctlParams.CopyIn(t, paramsAddr); err != nil { return 0, err } + numberOfPageTables := tpu.NumberOfTPUV4PageTables + if fd.device.lite { + numberOfPageTables = tpu.NumberOfTPUV4litePageTables + } + if userIoctlParams.PageTableIndex >= numberOfPageTables { + return 0, linuxerr.EFAULT + } + tmm := t.MemoryManager() ar, ok := tmm.CheckIORange(hostarch.Addr(userIoctlParams.HostAddress), int64(userIoctlParams.Size)) if !ok { @@ -45,6 +54,23 @@ func gasketMapBufferIoctl(ctx context.Context, t *kernel.Task, hostFd int32, fd if !ar.IsPageAligned() || (userIoctlParams.Size/hostarch.PageSize) == 0 { return 0, linuxerr.EINVAL } + + devAddr := userIoctlParams.DeviceAddress + // The kernel driver does not enforce page alignment on the device + // address although it will be implicitly rounded down to a page + // boundary. We do it explicitly because it simplifies tracking + // of allocated ranges in 'devAddrSet'. + devAddr &^= (hostarch.PageSize - 1) + + // Make sure that the device address range can be mapped. + devar := DevAddrRange{ + devAddr, + devAddr + userIoctlParams.Size, + } + if !devar.WellFormed() { + return 0, linuxerr.EINVAL + } + // Reserve a range in our address space. m, _, errno := unix.RawSyscall6(unix.SYS_MMAP, 0 /* addr */, uintptr(ar.Length()), unix.PROT_NONE, unix.MAP_PRIVATE|unix.MAP_ANONYMOUS, ^uintptr(0) /* fd */, 0 /* offset */) if errno != 0 { @@ -89,7 +115,6 @@ func gasketMapBufferIoctl(ctx context.Context, t *kernel.Task, hostFd int32, fd fd.device.mu.Lock() defer fd.device.mu.Unlock() - devAddr := userIoctlParams.DeviceAddress for _, pr := range prs { rlen := uint64(pr.Source.Length()) if !fd.device.devAddrSet.Add(DevAddrRange{ @@ -103,11 +128,30 @@ func gasketMapBufferIoctl(ctx context.Context, t *kernel.Task, hostFd int32, fd return n, nil } -func gasketUnmapBufferIoctl(ctx context.Context, t *kernel.Task, hostFd int32, fd *accelFD, paramsAddr hostarch.Addr) (uintptr, error) { +func gasketUnmapBufferIoctl(ctx context.Context, t *kernel.Task, hostFd int32, fd *tpuV4FD, paramsAddr hostarch.Addr) (uintptr, error) { var userIoctlParams gasket.GasketPageTableIoctl if _, err := userIoctlParams.CopyIn(t, paramsAddr); err != nil { return 0, err } + + numberOfPageTables := tpu.NumberOfTPUV4PageTables + if fd.device.lite { + numberOfPageTables = tpu.NumberOfTPUV4litePageTables + } + if userIoctlParams.PageTableIndex >= numberOfPageTables { + return 0, linuxerr.EFAULT + } + + devAddr := userIoctlParams.DeviceAddress + devAddr &^= (hostarch.PageSize - 1) + devar := DevAddrRange{ + devAddr, + devAddr + userIoctlParams.Size, + } + if !devar.WellFormed() { + return 0, linuxerr.EINVAL + } + sentryIoctlParams := userIoctlParams sentryIoctlParams.HostAddress = 0 // clobber this value, it's unused. n, err := ioctlInvokePtrArg(hostFd, gasket.GASKET_IOCTL_UNMAP_BUFFER, &sentryIoctlParams) @@ -129,12 +173,29 @@ func gasketUnmapBufferIoctl(ctx context.Context, t *kernel.Task, hostFd int32, f return n, nil } -func gasketInterruptMappingIoctl(ctx context.Context, t *kernel.Task, hostFd int32, paramsAddr hostarch.Addr) (uintptr, error) { +func gasketInterruptMappingIoctl(ctx context.Context, t *kernel.Task, hostFd int32, paramsAddr hostarch.Addr, lite bool) (uintptr, error) { var userIoctlParams gasket.GasketInterruptMapping if _, err := userIoctlParams.CopyIn(t, paramsAddr); err != nil { return 0, err } + sizeOfInterruptList := tpu.SizeOfTPUV4InterruptList + interruptMap := tpu.TPUV4InterruptsMap + if lite { + sizeOfInterruptList = tpu.SizeOfTPUV4liteInterruptList + interruptMap = tpu.TPUV4liteInterruptsMap + } + if userIoctlParams.Interrupt >= sizeOfInterruptList { + return 0, linuxerr.EINVAL + } + barRegMap, ok := interruptMap[userIoctlParams.BarIndex] + if !ok { + return 0, linuxerr.EINVAL + } + if _, ok := barRegMap[userIoctlParams.RegOffset]; !ok { + return 0, linuxerr.EINVAL + } + // Check that 'userEventFD.Eventfd' is an eventfd. eventFileGeneric, _ := t.FDTable().Get(int32(userIoctlParams.EventFD)) if eventFileGeneric == nil { diff --git a/pkg/sentry/devices/accel/accel_unsafe.go b/pkg/sentry/devices/accel/gasket_unsafe.go similarity index 100% rename from pkg/sentry/devices/accel/accel_unsafe.go rename to pkg/sentry/devices/accel/gasket_unsafe.go diff --git a/pkg/sentry/devices/accel/seccomp_filters.go b/pkg/sentry/devices/accel/seccomp_filters.go index 978fd4cb3..489d92cbd 100644 --- a/pkg/sentry/devices/accel/seccomp_filters.go +++ b/pkg/sentry/devices/accel/seccomp_filters.go @@ -41,30 +41,6 @@ func Filters() seccomp.SyscallRules { nonNegativeFD, seccomp.EqualTo(gasket.GASKET_IOCTL_RESET), }, - seccomp.PerArg{ - nonNegativeFD, - seccomp.EqualTo(gasket.GASKET_IOCTL_SET_EVENTFD), - }, - seccomp.PerArg{ - nonNegativeFD, - seccomp.EqualTo(gasket.GASKET_IOCTL_CLEAR_EVENTFD), - }, - seccomp.PerArg{ - nonNegativeFD, - seccomp.EqualTo(gasket.GASKET_IOCTL_NUMBER_PAGE_TABLES), - }, - seccomp.PerArg{ - nonNegativeFD, - seccomp.EqualTo(gasket.GASKET_IOCTL_PAGE_TABLE_SIZE), - }, - seccomp.PerArg{ - nonNegativeFD, - seccomp.EqualTo(gasket.GASKET_IOCTL_SIMPLE_PAGE_TABLE_SIZE), - }, - seccomp.PerArg{ - nonNegativeFD, - seccomp.EqualTo(gasket.GASKET_IOCTL_PARTITION_PAGE_TABLE), - }, seccomp.PerArg{ nonNegativeFD, seccomp.EqualTo(gasket.GASKET_IOCTL_MAP_BUFFER), @@ -85,10 +61,6 @@ func Filters() seccomp.SyscallRules { nonNegativeFD, seccomp.EqualTo(gasket.GASKET_IOCTL_UNREGISTER_INTERRUPT), }, - seccomp.PerArg{ - nonNegativeFD, - seccomp.EqualTo(gasket.GASKET_IOCTL_MAP_DMA_BUF), - }, }, unix.SYS_EVENTFD2: seccomp.Or{ seccomp.PerArg{ diff --git a/pkg/sentry/devices/accel/accel.go b/pkg/sentry/devices/accel/tpu_v4.go similarity index 85% rename from pkg/sentry/devices/accel/accel.go rename to pkg/sentry/devices/accel/tpu_v4.go index 6ce32158a..71c9f49cb 100644 --- a/pkg/sentry/devices/accel/accel.go +++ b/pkg/sentry/devices/accel/tpu_v4.go @@ -34,24 +34,24 @@ import ( "gvisor.dev/gvisor/pkg/waiter" ) -// accelFD implements vfs.FileDescriptionImpl for /dev/accel[0-9]+. +// tpuV4FD implements vfs.FileDescriptionImpl for /dev/accel[0-9]+. // // accelFD is not savable; we do not implement save/restore of accelerator // state. -type accelFD struct { +type tpuV4FD struct { vfsfd vfs.FileDescription vfs.FileDescriptionDefaultImpl vfs.DentryMetadataFileDescriptionImpl vfs.NoLockFD hostFD int32 - device *accelDevice + device *tpuV4Device queue waiter.Queue memmapFile accelFDMemmapFile } // Release implements vfs.FileDescriptionImpl.Release. -func (fd *accelFD) Release(context.Context) { +func (fd *tpuV4FD) Release(context.Context) { fd.device.mu.Lock() defer fd.device.mu.Unlock() fd.device.openWriteFDs-- @@ -75,13 +75,14 @@ func (fd *accelFD) Release(context.Context) { gap := s.Remove(seg) seg = gap.NextSegment() } + fd.device.owner = nil } fdnotifier.RemoveFD(fd.hostFD) unix.Close(int(fd.hostFD)) } // EventRegister implements waiter.Waitable.EventRegister. -func (fd *accelFD) EventRegister(e *waiter.Entry) error { +func (fd *tpuV4FD) EventRegister(e *waiter.Entry) error { fd.queue.EventRegister(e) if err := fdnotifier.UpdateFD(fd.hostFD); err != nil { fd.queue.EventUnregister(e) @@ -91,7 +92,7 @@ func (fd *accelFD) EventRegister(e *waiter.Entry) error { } // EventUnregister implements waiter.Waitable.EventUnregister. -func (fd *accelFD) EventUnregister(e *waiter.Entry) { +func (fd *tpuV4FD) EventUnregister(e *waiter.Entry) { fd.queue.EventUnregister(e) if err := fdnotifier.UpdateFD(fd.hostFD); err != nil { panic(fmt.Sprint("UpdateFD:", err)) @@ -99,17 +100,17 @@ func (fd *accelFD) EventUnregister(e *waiter.Entry) { } // Readiness implements waiter.Waitable.Readiness. -func (fd *accelFD) Readiness(mask waiter.EventMask) waiter.EventMask { +func (fd *tpuV4FD) Readiness(mask waiter.EventMask) waiter.EventMask { return fdnotifier.NonBlockingPoll(fd.hostFD, mask) } // Epollable implements vfs.FileDescriptionImpl.Epollable. -func (fd *accelFD) Epollable() bool { +func (fd *tpuV4FD) Epollable() bool { return true } // Ioctl implements vfs.FileDescriptionImpl.Ioctl. -func (fd *accelFD) Ioctl(ctx context.Context, uio usermem.IO, sysno uintptr, args arch.SyscallArguments) (uintptr, error) { +func (fd *tpuV4FD) Ioctl(ctx context.Context, uio usermem.IO, sysno uintptr, args arch.SyscallArguments) (uintptr, error) { cmd := args[1].Uint() argPtr := args[2].Pointer() argSize := linux.IOC_SIZE(cmd) @@ -118,6 +119,9 @@ func (fd *accelFD) Ioctl(ctx context.Context, uio usermem.IO, sysno uintptr, arg if t == nil { panic("Ioctl should be called from a task context") } + if err := fd.checkPermission(t); err != nil { + return 0, err + } log.Infof("Accel ioctl %s called on fd %d with arg %v of size %d.", gasket.Ioctl(cmd), fd.hostFD, argPtr, argSize) switch gasket.Ioctl(cmd) { @@ -136,7 +140,7 @@ func (fd *accelFD) Ioctl(ctx context.Context, uio usermem.IO, sysno uintptr, arg case gasket.GASKET_IOCTL_CLEAR_INTERRUPT_COUNTS: return ioctlInvoke(fd.hostFD, gasket.GASKET_IOCTL_CLEAR_INTERRUPT_COUNTS, 0) case gasket.GASKET_IOCTL_REGISTER_INTERRUPT: - return gasketInterruptMappingIoctl(ctx, t, fd.hostFD, argPtr) + return gasketInterruptMappingIoctl(ctx, t, fd.hostFD, argPtr, fd.device.lite) case gasket.GASKET_IOCTL_UNREGISTER_INTERRUPT: return ioctlInvoke[uint64](fd.hostFD, gasket.GASKET_IOCTL_UNREGISTER_INTERRUPT, args[2].Uint64()) default: @@ -144,6 +148,19 @@ func (fd *accelFD) Ioctl(ctx context.Context, uio usermem.IO, sysno uintptr, arg } } +// checkPermission checks that the thread that owns this device is the only +// one that can issue commands to the TPU. Other threads with access to +// /dev/accel will not be able to issue commands to the device. +func (fd *tpuV4FD) checkPermission(t *kernel.Task) error { + fd.device.mu.Lock() + defer fd.device.mu.Unlock() + owner := fd.device.owner + if t.ThreadGroup() != owner { + return linuxerr.EPERM + } + return nil +} + type pinnedAccelMem struct { pinnedRange mm.PinnedRange pageTableIndex uint64 diff --git a/pkg/sentry/devices/accel/accel_mmap.go b/pkg/sentry/devices/accel/tpu_v4_mmap.go similarity index 85% rename from pkg/sentry/devices/accel/accel_mmap.go rename to pkg/sentry/devices/accel/tpu_v4_mmap.go index ba55c9acd..77f222e66 100644 --- a/pkg/sentry/devices/accel/accel_mmap.go +++ b/pkg/sentry/devices/accel/tpu_v4_mmap.go @@ -25,26 +25,26 @@ import ( ) // ConfigureMMap implements vfs.FileDescriptionImpl.ConfigureMMap. -func (fd *accelFD) ConfigureMMap(ctx context.Context, opts *memmap.MMapOpts) error { +func (fd *tpuV4FD) ConfigureMMap(ctx context.Context, opts *memmap.MMapOpts) error { return vfs.GenericConfigureMMap(&fd.vfsfd, fd, opts) } // AddMapping implements memmap.Mappable.AddMapping. -func (fd *accelFD) AddMapping(ctx context.Context, ms memmap.MappingSpace, ar hostarch.AddrRange, offset uint64, writable bool) error { +func (fd *tpuV4FD) AddMapping(ctx context.Context, ms memmap.MappingSpace, ar hostarch.AddrRange, offset uint64, writable bool) error { return nil } // RemoveMapping implements memmap.Mappable.RemoveMapping. -func (fd *accelFD) RemoveMapping(ctx context.Context, ms memmap.MappingSpace, ar hostarch.AddrRange, offset uint64, writable bool) { +func (fd *tpuV4FD) RemoveMapping(ctx context.Context, ms memmap.MappingSpace, ar hostarch.AddrRange, offset uint64, writable bool) { } // CopyMapping implements memmap.Mappable.CopyMapping. -func (fd *accelFD) CopyMapping(ctx context.Context, ms memmap.MappingSpace, srcAR, dstAR hostarch.AddrRange, offset uint64, writable bool) error { +func (fd *tpuV4FD) CopyMapping(ctx context.Context, ms memmap.MappingSpace, srcAR, dstAR hostarch.AddrRange, offset uint64, writable bool) error { return nil } // Translate implements memmap.Mappable.Translate. -func (fd *accelFD) Translate(ctx context.Context, required, optional memmap.MappableRange, at hostarch.AccessType) ([]memmap.Translation, error) { +func (fd *tpuV4FD) Translate(ctx context.Context, required, optional memmap.MappableRange, at hostarch.AccessType) ([]memmap.Translation, error) { return []memmap.Translation{ { Source: optional, @@ -56,12 +56,12 @@ func (fd *accelFD) Translate(ctx context.Context, required, optional memmap.Mapp } // InvalidateUnsavable implements memmap.Mappable.InvalidateUnsavable. -func (fd *accelFD) InvalidateUnsavable(ctx context.Context) error { +func (fd *tpuV4FD) InvalidateUnsavable(ctx context.Context) error { return nil } type accelFDMemmapFile struct { - fd *accelFD + fd *tpuV4FD } // IncRef implements memmap.File.IncRef. diff --git a/runsc/boot/BUILD b/runsc/boot/BUILD index f2bc904f4..227737057 100644 --- a/runsc/boot/BUILD +++ b/runsc/boot/BUILD @@ -32,6 +32,7 @@ go_library( deps = [ "//pkg/abi", "//pkg/abi/linux", + "//pkg/abi/tpu", "//pkg/bpf", "//pkg/cleanup", "//pkg/context", diff --git a/runsc/boot/vfs.go b/runsc/boot/vfs.go index b453da45a..323a242e7 100644 --- a/runsc/boot/vfs.go +++ b/runsc/boot/vfs.go @@ -16,6 +16,7 @@ package boot import ( "fmt" + "os" "path" "path/filepath" "regexp" @@ -25,6 +26,7 @@ import ( specs "github.com/opencontainers/runtime-spec/specs-go" "gvisor.dev/gvisor/pkg/abi/linux" + "gvisor.dev/gvisor/pkg/abi/tpu" "gvisor.dev/gvisor/pkg/cleanup" "gvisor.dev/gvisor/pkg/context" "gvisor.dev/gvisor/pkg/errors/linuxerr" @@ -1168,11 +1170,32 @@ func tpuProxyRegisterDevicesAndCreateFiles(ctx context.Context, info *containerI if err != nil { return fmt.Errorf("enumerating accel device files: %w", err) } - for _, path := range paths { + pciAddrs, err := filepath.Glob("/sys/devices/pci0000:00/*") + if err != nil { + return fmt.Errorf("enumerating PCI device files: %w", err) + } + for _, accelPath := range paths { accelDeviceRegex := regexp.MustCompile(`^/dev/accel(\d+)$`) - if ms := accelDeviceRegex.FindStringSubmatch(path); ms != nil { + if ms := accelDeviceRegex.FindStringSubmatch(accelPath); ms != nil { deviceNum, _ := strconv.ParseUint(ms[1], 10, 32) - if err := accel.Register(vfsObj, uint32(deviceNum)); err != nil { + + var pciDevicePath string + for _, pciPath := range pciAddrs { + if _, err := os.Stat(path.Join(pciPath, fmt.Sprintf("accel/accel%d", deviceNum))); err == nil { + pciDevicePath = pciPath + } + } + var deviceIDBytes []byte + if deviceIDBytes, err = os.ReadFile(path.Join(pciDevicePath, "device")); err != nil { + return fmt.Errorf("reading PCI device ID: %w", err) + } + deviceIDStr := strings.Replace(string(deviceIDBytes), "0x", "", -1) + deviceID, err := strconv.ParseInt(strings.TrimSpace(deviceIDStr), 16, 64) + if err != nil { + return fmt.Errorf("parsing PCI device ID: %w", err) + } + + if err := accel.RegisterTPUV4Device(vfsObj, uint32(deviceNum), deviceID == tpu.TPUV4liteDeviceID); err != nil { return fmt.Errorf("registering accel driver: %w", err) } if err := accel.CreateDevtmpfsFile(ctx, a, uint32(deviceNum)); err != nil { diff --git a/runsc/cmd/util/BUILD b/runsc/cmd/util/BUILD index 8f11ae494..2d55017dd 100644 --- a/runsc/cmd/util/BUILD +++ b/runsc/cmd/util/BUILD @@ -17,6 +17,7 @@ go_library( "//tools:__subpackages__", ], deps = [ + "//pkg/abi/tpu", "//pkg/log", "@com_github_google_subcommands//:go_default_library", ], diff --git a/runsc/cmd/util/tpu.go b/runsc/cmd/util/tpu.go index 1a4c9d93a..3f11af8a9 100644 --- a/runsc/cmd/util/tpu.go +++ b/runsc/cmd/util/tpu.go @@ -21,11 +21,13 @@ import ( "regexp" "strconv" "strings" + + "gvisor.dev/gvisor/pkg/abi/tpu" ) const googleVendorID = 0x1AE0 -var tpuV4DeviceIDs = map[uint64]any{0x005E: nil, 0x0056: nil} +var tpuV4DeviceIDs = map[uint64]any{tpu.TPUV4DeviceID: nil, tpu.TPUV4liteDeviceID: nil} // TODO(b/288456802): Add support for /dev/vfio controlled accelerators. // This is required for v5+ TPUs.