Implement ioctl command VFIO_IOMMU_MAP_DMA.

PiperOrigin-RevId: 619691808
This commit is contained in:
Jing Chen
2024-03-27 16:10:41 -07:00
committed by gVisor bot
parent db85b6316f
commit 79dd2520ff
6 changed files with 226 additions and 12 deletions
+26
View File
@@ -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
}
+33
View File
@@ -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",
},
)
+3
View File
@@ -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.
+24 -12
View File
@@ -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),
},
},
})
}
+55
View File
@@ -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
}
+85
View File
@@ -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