From 79dd2520ffd8ea1213662f4df11b3eb709ae23c9 Mon Sep 17 00:00:00 2001 From: Jing Chen Date: Wed, 27 Mar 2024 16:07:09 -0700 Subject: [PATCH] Implement ioctl command VFIO_IOMMU_MAP_DMA. PiperOrigin-RevId: 619691808 --- pkg/abi/linux/vfio.go | 26 ++++++ pkg/sentry/devices/tpuproxy/BUILD | 33 +++++++ pkg/sentry/devices/tpuproxy/device.go | 3 + pkg/sentry/devices/tpuproxy/seccomp_filter.go | 36 +++++--- pkg/sentry/devices/tpuproxy/tpu.go | 55 ++++++++++++ pkg/sentry/devices/tpuproxy/vfio.go | 85 +++++++++++++++++++ 6 files changed, 226 insertions(+), 12 deletions(-) diff --git a/pkg/abi/linux/vfio.go b/pkg/abi/linux/vfio.go index 44f458899..56c1101e5 100644 --- a/pkg/abi/linux/vfio.go +++ b/pkg/abi/linux/vfio.go @@ -96,6 +96,16 @@ const ( VFIO_PCI_NUM_IRQS ) +// VFIOIommuType1DmaMap flags. +const ( + // Readable from device. + VFIO_DMA_MAP_FLAG_READ = 1 << iota + // Writable from device. + VFIO_DMA_MAP_FLAG_WRITE + // Update the device's virtual address. + VFIO_DMA_MAP_FLAG_VADDR +) + // IOCTLs for VFIO file descriptor from include/uapi/linux/vfio.h. var ( VFIO_CHECK_EXTENSION = IO(VFIO_TYPE, VFIO_BASE+1) @@ -107,6 +117,7 @@ var ( VFIO_DEVICE_GET_IRQ_INFO = IO(VFIO_TYPE, VFIO_BASE+9) VFIO_DEVICE_SET_IRQS = IO(VFIO_TYPE, VFIO_BASE+10) VFIO_DEVICE_RESET = IO(VFIO_TYPE, VFIO_BASE+11) + VFIO_IOMMU_MAP_DMA = IO(VFIO_TYPE, VFIO_BASE+13) ) // VFIODeviceInfo is analogous to vfio_device_info @@ -165,3 +176,18 @@ type VFIOIrqSet struct { Start uint32 Count uint32 } + +// VFIOIommuType1DmaMap is analogous to vfio_iommu_type1_dma_map +// from include/uapi/linux/vfio.h. +// +// +marshal +type VFIOIommuType1DmaMap struct { + Argsz uint32 + Flags uint32 + // Process virtual address. + Vaddr uint64 + // IO virtual address. + IOVa uint64 + // Size of mapping in bytes. + Size uint64 +} diff --git a/pkg/sentry/devices/tpuproxy/BUILD b/pkg/sentry/devices/tpuproxy/BUILD index 39e2b5f1e..f85405fef 100644 --- a/pkg/sentry/devices/tpuproxy/BUILD +++ b/pkg/sentry/devices/tpuproxy/BUILD @@ -1,4 +1,5 @@ load("//tools:defs.bzl", "go_library") +load("//tools/go_generics:defs.bzl", "go_template_instance") package(default_applicable_licenses = ["//:license"]) @@ -7,6 +8,8 @@ licenses(["notice"]) go_library( name = "tpuproxy", srcs = [ + "devaddr_range.go", + "devaddr_set.go", "device.go", "ioctl_unsafe.go", "seccomp_filter.go", @@ -20,6 +23,7 @@ go_library( ], deps = [ "//pkg/abi/linux", + "//pkg/cleanup", "//pkg/context", "//pkg/devutil", "//pkg/errors/linuxerr", @@ -34,6 +38,7 @@ go_library( "//pkg/sentry/fsimpl/kernfs", "//pkg/sentry/kernel", "//pkg/sentry/memmap", + "//pkg/sentry/mm", "//pkg/sentry/vfs", "//pkg/sync", "//pkg/usermem", @@ -42,3 +47,31 @@ go_library( "@org_golang_x_sys//unix:go_default_library", ], ) + +go_template_instance( + name = "devaddr_range", + out = "devaddr_range.go", + package = "tpuproxy", + prefix = "DevAddr", + template = "//pkg/segment:generic_range", + types = { + "T": "uint64", + }, +) + +go_template_instance( + name = "devaddr_set", + out = "devaddr_set.go", + imports = { + "mm": "gvisor.dev/gvisor/pkg/sentry/mm", + }, + package = "tpuproxy", + prefix = "DevAddr", + template = "//pkg/segment:generic_set", + types = { + "Key": "uint64", + "Range": "DevAddrRange", + "Value": "mm.PinnedRange", + "Functions": "devAddrSetFuncs", + }, +) diff --git a/pkg/sentry/devices/tpuproxy/device.go b/pkg/sentry/devices/tpuproxy/device.go index 667a8b9a6..2f1245c26 100644 --- a/pkg/sentry/devices/tpuproxy/device.go +++ b/pkg/sentry/devices/tpuproxy/device.go @@ -85,6 +85,9 @@ func (dev *tpuDevice) Open(ctx context.Context, mnt *vfs.Mount, d *vfs.Dentry, o // device implements vfs.Device for /dev/vfio/vfio. type vfioDevice struct { mu sync.Mutex + + // +checklocks:mu + devAddrSet DevAddrSet } // Open implements vfs.Device.Open. diff --git a/pkg/sentry/devices/tpuproxy/seccomp_filter.go b/pkg/sentry/devices/tpuproxy/seccomp_filter.go index 18bd2a299..f4f515d97 100644 --- a/pkg/sentry/devices/tpuproxy/seccomp_filter.go +++ b/pkg/sentry/devices/tpuproxy/seccomp_filter.go @@ -52,25 +52,21 @@ func Filters() seccomp.SyscallRules { seccomp.AnyValue{}, seccomp.EqualTo(0), }, + unix.SYS_MMAP: seccomp.PerArg{ + seccomp.AnyValue{}, + seccomp.AnyValue{}, + seccomp.EqualTo(linux.PROT_READ | linux.PROT_WRITE), + seccomp.EqualTo(linux.MAP_SHARED | linux.MAP_LOCKED), + seccomp.NonNegativeFD{}, + }, + unix.SYS_MUNMAP: seccomp.MatchAll{}, unix.SYS_PREAD64: seccomp.MatchAll{}, unix.SYS_PWRITE64: seccomp.MatchAll{}, unix.SYS_IOCTL: seccomp.Or{ - seccomp.PerArg{ - seccomp.NonNegativeFD{}, - seccomp.EqualTo(linux.VFIO_GROUP_SET_CONTAINER), - }, seccomp.PerArg{ seccomp.NonNegativeFD{}, seccomp.EqualTo(linux.VFIO_CHECK_EXTENSION), }, - seccomp.PerArg{ - seccomp.NonNegativeFD{}, - seccomp.EqualTo(linux.VFIO_SET_IOMMU), - }, - seccomp.PerArg{ - seccomp.NonNegativeFD{}, - seccomp.EqualTo(linux.VFIO_GROUP_GET_DEVICE_FD), - }, seccomp.PerArg{ seccomp.NonNegativeFD{}, seccomp.EqualTo(linux.VFIO_DEVICE_GET_INFO), @@ -87,6 +83,22 @@ func Filters() seccomp.SyscallRules { seccomp.NonNegativeFD{}, seccomp.EqualTo(linux.VFIO_DEVICE_SET_IRQS), }, + seccomp.PerArg{ + seccomp.NonNegativeFD{}, + seccomp.EqualTo(linux.VFIO_GROUP_GET_DEVICE_FD), + }, + seccomp.PerArg{ + seccomp.NonNegativeFD{}, + seccomp.EqualTo(linux.VFIO_GROUP_SET_CONTAINER), + }, + seccomp.PerArg{ + seccomp.NonNegativeFD{}, + seccomp.EqualTo(linux.VFIO_IOMMU_MAP_DMA), + }, + seccomp.PerArg{ + seccomp.NonNegativeFD{}, + seccomp.EqualTo(linux.VFIO_SET_IOMMU), + }, }, }) } diff --git a/pkg/sentry/devices/tpuproxy/tpu.go b/pkg/sentry/devices/tpuproxy/tpu.go index 9d9e2c1cd..381627102 100644 --- a/pkg/sentry/devices/tpuproxy/tpu.go +++ b/pkg/sentry/devices/tpuproxy/tpu.go @@ -29,6 +29,7 @@ import ( "gvisor.dev/gvisor/pkg/sentry/fsimpl/eventfd" "gvisor.dev/gvisor/pkg/sentry/fsimpl/kernfs" "gvisor.dev/gvisor/pkg/sentry/kernel" + "gvisor.dev/gvisor/pkg/sentry/mm" "gvisor.dev/gvisor/pkg/sentry/vfs" "gvisor.dev/gvisor/pkg/usermem" "gvisor.dev/gvisor/pkg/waiter" @@ -424,3 +425,57 @@ func (fd *pciDeviceFD) PWrite(ctx context.Context, src usermem.IOSequence, offse n, err := unix.Pwrite(int(fd.hostFD), buf, offset) return int64(n), err } + +// DevAddrSet tracks device address ranges that have been mapped. +type devAddrSetFuncs struct{} + +func (devAddrSetFuncs) MinKey() uint64 { + return 0 +} + +func (devAddrSetFuncs) MaxKey() uint64 { + return ^uint64(0) +} + +func (devAddrSetFuncs) ClearValue(val *mm.PinnedRange) { + *val = mm.PinnedRange{} +} + +func (devAddrSetFuncs) Merge(r1 DevAddrRange, v1 mm.PinnedRange, r2 DevAddrRange, v2 mm.PinnedRange) (mm.PinnedRange, bool) { + // Do we have the same backing file? + if v1.File != v2.File { + return mm.PinnedRange{}, false + } + + // Do we have contiguous offsets in the backing file? + if v1.Offset+uint64(v1.Source.Length()) != v2.Offset { + return mm.PinnedRange{}, false + } + + // Are the virtual addresses contiguous? + // + // This check isn't strictly needed because 'mm.PinnedRange.Source' + // is only used to track the size of the pinned region (this is + // because the virtual address range can be unmapped or remapped + // elsewhere). Regardless we require this for simplicity. + if v1.Source.End != v2.Source.Start { + return mm.PinnedRange{}, false + } + + // Extend v1 to account for the adjacent PinnedRange. + v1.Source.End = v2.Source.End + return v1, true +} + +func (devAddrSetFuncs) Split(r DevAddrRange, val mm.PinnedRange, split uint64) (mm.PinnedRange, mm.PinnedRange) { + n := split - r.Start + + left := val + left.Source.End = left.Source.Start + hostarch.Addr(n) + + right := val + right.Source.Start += hostarch.Addr(n) + right.Offset += n + + return left, right +} diff --git a/pkg/sentry/devices/tpuproxy/vfio.go b/pkg/sentry/devices/tpuproxy/vfio.go index fd47550de..eb809741d 100644 --- a/pkg/sentry/devices/tpuproxy/vfio.go +++ b/pkg/sentry/devices/tpuproxy/vfio.go @@ -19,12 +19,16 @@ import ( "golang.org/x/sys/unix" "gvisor.dev/gvisor/pkg/abi/linux" + "gvisor.dev/gvisor/pkg/cleanup" "gvisor.dev/gvisor/pkg/context" "gvisor.dev/gvisor/pkg/errors/linuxerr" "gvisor.dev/gvisor/pkg/fdnotifier" + "gvisor.dev/gvisor/pkg/hostarch" "gvisor.dev/gvisor/pkg/log" "gvisor.dev/gvisor/pkg/sentry/arch" "gvisor.dev/gvisor/pkg/sentry/kernel" + "gvisor.dev/gvisor/pkg/sentry/memmap" + "gvisor.dev/gvisor/pkg/sentry/mm" "gvisor.dev/gvisor/pkg/sentry/vfs" "gvisor.dev/gvisor/pkg/usermem" "gvisor.dev/gvisor/pkg/waiter" @@ -90,6 +94,8 @@ func (fd *vfioFd) Ioctl(ctx context.Context, uio usermem.IO, sysno uintptr, args return fd.checkExtension(extension(args[2].Int())) case linux.VFIO_SET_IOMMU: return fd.setIOMMU(extension(args[2].Int())) + case linux.VFIO_IOMMU_MAP_DMA: + return fd.iommuMapDma(ctx, t, args[2].Pointer()) } return 0, linuxerr.ENOSYS } @@ -124,6 +130,85 @@ func (fd *vfioFd) setIOMMU(ext extension) (uintptr, error) { return 0, linuxerr.EINVAL } +func (fd *vfioFd) iommuMapDma(ctx context.Context, t *kernel.Task, arg hostarch.Addr) (uintptr, error) { + var dmaMap linux.VFIOIommuType1DmaMap + if _, err := dmaMap.CopyIn(t, arg); err != nil { + return 0, err + } + tmm := t.MemoryManager() + ar, ok := tmm.CheckIORange(hostarch.Addr(dmaMap.Vaddr), int64(dmaMap.Size)) + if !ok { + return 0, linuxerr.EFAULT + } + if !ar.IsPageAligned() || (dmaMap.Size/hostarch.PageSize) == 0 { + return 0, linuxerr.EINVAL + } + // See comments at pkg/sentry/devices/accel/gasket.go, line 57-60. + devAddr := dmaMap.IOVa + devAddr &^= (hostarch.PageSize - 1) + + devar := DevAddrRange{ + devAddr, + devAddr + dmaMap.Size, + } + if !devar.WellFormed() { + return 0, linuxerr.EINVAL + } + // Reserve a range in the 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), 0) + if errno != 0 { + return 0, errno + } + cu := cleanup.Make(func() { + unix.RawSyscall(unix.SYS_MUNMAP, m, uintptr(ar.Length()), 0) + }) + defer cu.Clean() + // Mirror application mappings into the reserved range. + prs, err := t.MemoryManager().Pin(ctx, ar, hostarch.ReadWrite, false) + cu.Add(func() { + mm.Unpin(prs) + }) + if err != nil { + return 0, err + } + sentryAddr := uintptr(m) + for _, pr := range prs { + ims, err := pr.File.MapInternal(memmap.FileRange{Start: pr.Offset, End: pr.Offset + uint64(pr.Source.Length())}, hostarch.ReadWrite) + if err != nil { + return 0, err + } + for !ims.IsEmpty() { + im := ims.Head() + if _, _, errno := unix.RawSyscall6(unix.SYS_MREMAP, im.Addr(), 0, uintptr(im.Len()), linux.MREMAP_MAYMOVE|linux.MREMAP_FIXED, sentryAddr, 0); errno != 0 { + return 0, errno + } + sentryAddr += uintptr(im.Len()) + ims = ims.Tail() + } + } + // Replace Vaddr with the host's virtual address. + dmaMap.Vaddr = uint64(m) + n, err := IOCTLInvokePtrArg[uint32](fd.hostFd, linux.VFIO_IOMMU_MAP_DMA, &dmaMap) + if err != nil { + return n, err + } + cu.Release() + // Unmap the reserved range, which is no longer required. + unix.RawSyscall(unix.SYS_MUNMAP, m, uintptr(ar.Length()), 0) + + fd.device.mu.Lock() + defer fd.device.mu.Unlock() + for _, pr := range prs { + rlen := uint64(pr.Source.Length()) + fd.device.devAddrSet.InsertRange(DevAddrRange{ + devAddr, + devAddr + rlen, + }, pr) + devAddr += rlen + } + return n, nil +} + // VFIO extension. type extension int32