mirror of
https://github.com/netbirdio/gvisor.git
synced 2026-05-22 17:12:49 -07:00
Implement ioctl command VFIO_IOMMU_MAP_DMA.
PiperOrigin-RevId: 619691808
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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",
|
||||
},
|
||||
)
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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),
|
||||
},
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user