From 1d4792c56655dcd8099babe622431abbe444ecd2 Mon Sep 17 00:00:00 2001 From: Lucas Manning Date: Mon, 31 Jul 2023 18:03:06 -0700 Subject: [PATCH] Implement accel fd methods and gasket ioctls. The implementation of memory mapping and interrupt registration is very similar to what's already been done for nvproxy. PiperOrigin-RevId: 552644264 --- WORKSPACE | 14 +- go.mod | 3 +- go.sum | 4 + pkg/sentry/devices/accel/BUILD | 45 ++++++ pkg/sentry/devices/accel/accel.go | 144 ++++++++++++++++++- pkg/sentry/devices/accel/accel_mmap.go | 84 ++++++++++++ pkg/sentry/devices/accel/accel_unsafe.go | 35 +++++ pkg/sentry/devices/accel/device.go | 19 +++ pkg/sentry/devices/accel/gasket.go | 167 +++++++++++++++++++++++ 9 files changed, 504 insertions(+), 11 deletions(-) create mode 100644 pkg/sentry/devices/accel/accel_mmap.go create mode 100644 pkg/sentry/devices/accel/accel_unsafe.go create mode 100644 pkg/sentry/devices/accel/gasket.go diff --git a/WORKSPACE b/WORKSPACE index 7476de4f9..fb79487c2 100644 --- a/WORKSPACE +++ b/WORKSPACE @@ -102,6 +102,13 @@ go_repository( version = "v0.4.0", ) +go_repository( + name = "org_golang_x_exp", + importpath = "golang.org/x/exp", + sum = "h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY=", + version = "v0.0.0-20230725093048-515e97ebf090", +) + go_repository( name = "org_golang_x_net", importpath = "golang.org/x/net", @@ -2297,13 +2304,6 @@ go_repository( version = "v4.2.3", ) -go_repository( - name = "org_golang_x_exp", - importpath = "golang.org/x/exp", - sum = "h1:c2HOrn5iMezYjSlGPncknSEr/8x5LELb/ilJbXi9DEA=", - version = "v0.0.0-20190121172915-509febef88a4", -) - go_repository( name = "org_uber_go_tools", importpath = "go.uber.org/tools", diff --git a/go.mod b/go.mod index bcc3d6811..f95dfc1d4 100644 --- a/go.mod +++ b/go.mod @@ -26,7 +26,7 @@ require ( github.com/sirupsen/logrus v1.8.1 github.com/syndtr/gocapability v0.0.0-20200815063812-42c35b437635 github.com/vishvananda/netlink v1.1.1-0.20211118161826-650dca95af54 - golang.org/x/mod v0.7.0 + golang.org/x/mod v0.11.0 golang.org/x/sync v0.1.0 golang.org/x/sys v0.4.0 golang.org/x/time v0.0.0-20220210224613-90d013bbcef8 @@ -60,6 +60,7 @@ require ( github.com/pkg/errors v0.9.1 // indirect github.com/vishvananda/netns v0.0.0-20200728191858-db3c7e526aae // indirect go.opencensus.io v0.24.0 // indirect + golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 // indirect golang.org/x/net v0.5.0 // indirect golang.org/x/oauth2 v0.4.0 // indirect golang.org/x/term v0.4.0 // indirect diff --git a/go.sum b/go.sum index a7f95f9cd..39ee2c11f 100644 --- a/go.sum +++ b/go.sum @@ -233,6 +233,8 @@ golang.org/x/crypto v0.0.0-20190308221718-c2843e01d9a2/go.mod h1:djNgcEr1/C05ACk golang.org/x/crypto v0.0.0-20191011191535-87dc89f01550/go.mod h1:yigFU9vqHzYiE8UmvKecakEJjdnWj3jj499lnFckfCI= golang.org/x/crypto v0.0.0-20200622213623-75b288015ac9/go.mod h1:LzIPMQfyMNhhGPhUkYOs5KpL4U8rLKemX1yGLhDgUto= golang.org/x/exp v0.0.0-20190121172915-509febef88a4/go.mod h1:CJ0aWSM057203Lf6IL+f9T1iT9GByDxfZKAQTCR3kQA= +golang.org/x/exp v0.0.0-20230725093048-515e97ebf090 h1:Di6/M8l0O2lCLc6VVRWhgCiApHV8MnQurBnFSHsQtNY= +golang.org/x/exp v0.0.0-20230725093048-515e97ebf090/go.mod h1:FXUEEKJgO7OQYeo8N01OfiKP8RXMtf6e8aTskBGqWdc= golang.org/x/lint v0.0.0-20181026193005-c67002cb31c3/go.mod h1:UVdnD1Gm6xHRNCYTkRU2/jEulfH38KcIWyp/GAMgvoE= golang.org/x/lint v0.0.0-20190227174305-5b3e6a55c961/go.mod h1:wehouNa3lNwaWXcvxsM5YxQ5yQlVC4a0KAMCusXpPoU= golang.org/x/lint v0.0.0-20190313153728-d0100b6bd8b3/go.mod h1:6SW0HCj/g11FgYtHlgUYUwCkIfeOF89ocIRzGO/8vkc= @@ -241,6 +243,8 @@ golang.org/x/mod v0.3.0/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.4.2/go.mod h1:s0Qsj1ACt9ePp/hMypM3fl4fZqREWJwdYDEqhRiZZUA= golang.org/x/mod v0.7.0 h1:LapD9S96VoQRhi/GrNTqeBJFrUjs5UHCAtTlgwA5oZA= golang.org/x/mod v0.7.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs= +golang.org/x/mod v0.11.0 h1:bUO06HqtnRcc/7l71XBe4WcqTZ+3AH1J59zWDDwLKgU= +golang.org/x/mod v0.11.0/go.mod h1:iBbtSCu2XBx23ZKBPSOrRkjjQPZFPuis4dIYUhu/chs= golang.org/x/net v0.0.0-20180724234803-3673e40ba225/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20180826012351-8a410e7b638d/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= golang.org/x/net v0.0.0-20190213061140-3a22650c66bd/go.mod h1:mL1N/T3taQHkDXs73rZJwtUhF3w3ftmwwsq0BUmARs4= diff --git a/pkg/sentry/devices/accel/BUILD b/pkg/sentry/devices/accel/BUILD index 13a983d96..29a9c2388 100644 --- a/pkg/sentry/devices/accel/BUILD +++ b/pkg/sentry/devices/accel/BUILD @@ -1,3 +1,4 @@ +load("//tools/go_generics:defs.bzl", "go_template_instance") load("//tools:defs.bzl", "go_library") licenses(["notice"]) @@ -6,21 +7,65 @@ go_library( name = "accel", srcs = [ "accel.go", + "accel_mmap.go", + "accel_unsafe.go", + "devaddr_range.go", + "devaddr_set.go", "device.go", + "gasket.go", "seccomp_filters.go", ], visibility = ["//pkg/sentry:internal"], deps = [ "//pkg/abi/gasket", "//pkg/abi/linux", + "//pkg/cleanup", "//pkg/context", "//pkg/errors/linuxerr", + "//pkg/fdnotifier", + "//pkg/hostarch", + "//pkg/log", + "//pkg/safemem", "//pkg/seccomp", "//pkg/sentry/arch", "//pkg/sentry/fsimpl/devtmpfs", + "//pkg/sentry/fsimpl/eventfd", + "//pkg/sentry/kernel", + "//pkg/sentry/memmap", + "//pkg/sentry/mm", "//pkg/sentry/vfs", + "//pkg/sync", "//pkg/usermem", "//pkg/waiter", + "@org_golang_x_exp//constraints:go_default_library", "@org_golang_x_sys//unix:go_default_library", ], ) + +go_template_instance( + name = "devaddr_range", + out = "devaddr_range.go", + package = "accel", + 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 = "accel", + prefix = "DevAddr", + template = "//pkg/segment:generic_set", + types = { + "Key": "uint64", + "Range": "DevAddrRange", + "Value": "pinnedAccelMem", + "Functions": "devAddrSetFuncs", + }, +) diff --git a/pkg/sentry/devices/accel/accel.go b/pkg/sentry/devices/accel/accel.go index 97868dc63..6ce32158a 100644 --- a/pkg/sentry/devices/accel/accel.go +++ b/pkg/sentry/devices/accel/accel.go @@ -16,9 +16,19 @@ package accel import ( + "fmt" + + "golang.org/x/sys/unix" + "gvisor.dev/gvisor/pkg/abi/gasket" + "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/hostarch" + "gvisor.dev/gvisor/pkg/log" "gvisor.dev/gvisor/pkg/sentry/arch" + "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" @@ -34,25 +44,63 @@ type accelFD struct { vfs.DentryMetadataFileDescriptionImpl vfs.NoLockFD - hostFD int32 + hostFD int32 + device *accelDevice + queue waiter.Queue + memmapFile accelFDMemmapFile } // Release implements vfs.FileDescriptionImpl.Release. func (fd *accelFD) Release(context.Context) { + fd.device.mu.Lock() + defer fd.device.mu.Unlock() + fd.device.openWriteFDs-- + if fd.device.openWriteFDs == 0 { + log.Infof("openWriteFDs is zero, unpinning all sentry memory mappings") + s := &fd.device.devAddrSet + seg := s.FirstSegment() + for seg.Ok() { + r, v := seg.Range(), seg.Value() + gpti := gasket.GasketPageTableIoctl{ + PageTableIndex: v.pageTableIndex, + DeviceAddress: r.Start, + Size: r.End - r.Start, + HostAddress: 0, + } + _, err := ioctlInvokePtrArg(fd.hostFD, gasket.GASKET_IOCTL_UNMAP_BUFFER, &gpti) + if err != nil { + log.Warningf("could not unmap range [%#x, %#x) (index %d) on device: %v", r.Start, r.End, v.pageTableIndex, err) + } + mm.Unpin([]mm.PinnedRange{v.pinnedRange}) + gap := s.Remove(seg) + seg = gap.NextSegment() + } + } + fdnotifier.RemoveFD(fd.hostFD) + unix.Close(int(fd.hostFD)) } // EventRegister implements waiter.Waitable.EventRegister. func (fd *accelFD) EventRegister(e *waiter.Entry) error { + fd.queue.EventRegister(e) + if err := fdnotifier.UpdateFD(fd.hostFD); err != nil { + fd.queue.EventUnregister(e) + return err + } return nil } // EventUnregister implements waiter.Waitable.EventUnregister. func (fd *accelFD) EventUnregister(e *waiter.Entry) { + fd.queue.EventUnregister(e) + if err := fdnotifier.UpdateFD(fd.hostFD); err != nil { + panic(fmt.Sprint("UpdateFD:", err)) + } } // Readiness implements waiter.Waitable.Readiness. func (fd *accelFD) Readiness(mask waiter.EventMask) waiter.EventMask { - return waiter.EventErr + return fdnotifier.NonBlockingPoll(fd.hostFD, mask) } // Epollable implements vfs.FileDescriptionImpl.Epollable. @@ -62,5 +110,95 @@ func (fd *accelFD) Epollable() bool { // Ioctl implements vfs.FileDescriptionImpl.Ioctl. func (fd *accelFD) Ioctl(ctx context.Context, uio usermem.IO, sysno uintptr, args arch.SyscallArguments) (uintptr, error) { - return 0, linuxerr.ENOSYS + cmd := args[1].Uint() + argPtr := args[2].Pointer() + argSize := linux.IOC_SIZE(cmd) + + t := kernel.TaskFromContext(ctx) + if t == nil { + panic("Ioctl should be called from a task context") + } + + 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) { + // Not yet implemented gasket ioctls. + case gasket.GASKET_IOCTL_SET_EVENTFD, gasket.GASKET_IOCTL_CLEAR_EVENTFD, + gasket.GASKET_IOCTL_NUMBER_PAGE_TABLES, gasket.GASKET_IOCTL_PAGE_TABLE_SIZE, + gasket.GASKET_IOCTL_SIMPLE_PAGE_TABLE_SIZE, gasket.GASKET_IOCTL_PARTITION_PAGE_TABLE, + gasket.GASKET_IOCTL_MAP_DMA_BUF: + return 0, linuxerr.ENOSYS + case gasket.GASKET_IOCTL_RESET: + return ioctlInvoke[uint64](fd.hostFD, gasket.GASKET_IOCTL_RESET, args[2].Uint64()) + case gasket.GASKET_IOCTL_MAP_BUFFER: + return gasketMapBufferIoctl(ctx, t, fd.hostFD, fd, argPtr) + case gasket.GASKET_IOCTL_UNMAP_BUFFER: + return gasketUnmapBufferIoctl(ctx, t, fd.hostFD, fd, argPtr) + 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) + case gasket.GASKET_IOCTL_UNREGISTER_INTERRUPT: + return ioctlInvoke[uint64](fd.hostFD, gasket.GASKET_IOCTL_UNREGISTER_INTERRUPT, args[2].Uint64()) + default: + return 0, linuxerr.EINVAL + } +} + +type pinnedAccelMem struct { + pinnedRange mm.PinnedRange + pageTableIndex uint64 +} + +// 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 *pinnedAccelMem) { + *val = pinnedAccelMem{} +} + +func (devAddrSetFuncs) Merge(r1 DevAddrRange, v1 pinnedAccelMem, r2 DevAddrRange, v2 pinnedAccelMem) (pinnedAccelMem, bool) { + // Do we have the same backing file? + if v1.pinnedRange.File != v2.pinnedRange.File { + return pinnedAccelMem{}, false + } + + // Do we have contiguous offsets in the backing file? + if v1.pinnedRange.Offset+uint64(v1.pinnedRange.Source.Length()) != v2.pinnedRange.Offset { + return pinnedAccelMem{}, 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.pinnedRange.Source.End != v2.pinnedRange.Source.Start { + return pinnedAccelMem{}, false + } + + // Extend v1 to account for the adjacent PinnedRange. + v1.pinnedRange.Source.End = v2.pinnedRange.Source.End + return v1, true +} + +func (devAddrSetFuncs) Split(r DevAddrRange, val pinnedAccelMem, split uint64) (pinnedAccelMem, pinnedAccelMem) { + n := split - r.Start + + left := val + left.pinnedRange.Source.End = left.pinnedRange.Source.Start + hostarch.Addr(n) + + right := val + right.pinnedRange.Source.Start += hostarch.Addr(n) + right.pinnedRange.Offset += n + + return left, right } diff --git a/pkg/sentry/devices/accel/accel_mmap.go b/pkg/sentry/devices/accel/accel_mmap.go new file mode 100644 index 000000000..ba55c9acd --- /dev/null +++ b/pkg/sentry/devices/accel/accel_mmap.go @@ -0,0 +1,84 @@ +// 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 accel + +import ( + "gvisor.dev/gvisor/pkg/context" + "gvisor.dev/gvisor/pkg/errors/linuxerr" + "gvisor.dev/gvisor/pkg/hostarch" + "gvisor.dev/gvisor/pkg/log" + "gvisor.dev/gvisor/pkg/safemem" + "gvisor.dev/gvisor/pkg/sentry/memmap" + "gvisor.dev/gvisor/pkg/sentry/vfs" +) + +// ConfigureMMap implements vfs.FileDescriptionImpl.ConfigureMMap. +func (fd *accelFD) 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 { + return nil +} + +// RemoveMapping implements memmap.Mappable.RemoveMapping. +func (fd *accelFD) 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 { + return nil +} + +// Translate implements memmap.Mappable.Translate. +func (fd *accelFD) Translate(ctx context.Context, required, optional memmap.MappableRange, at hostarch.AccessType) ([]memmap.Translation, error) { + return []memmap.Translation{ + { + Source: optional, + File: &fd.memmapFile, + Offset: optional.Start, + Perms: at, + }, + }, nil +} + +// InvalidateUnsavable implements memmap.Mappable.InvalidateUnsavable. +func (fd *accelFD) InvalidateUnsavable(ctx context.Context) error { + return nil +} + +type accelFDMemmapFile struct { + fd *accelFD +} + +// IncRef implements memmap.File.IncRef. +func (mf *accelFDMemmapFile) IncRef(memmap.FileRange, uint32) { +} + +// DecRef implements memmap.File.DecRef. +func (mf *accelFDMemmapFile) DecRef(fr memmap.FileRange) { +} + +// MapInternal implements memmap.File.MapInternal. +func (mf *accelFDMemmapFile) MapInternal(fr memmap.FileRange, at hostarch.AccessType) (safemem.BlockSeq, error) { + log.Traceback("accel: rejecting accelFDMemmapFile.MapInternal") + return safemem.BlockSeq{}, linuxerr.EINVAL +} + +// FD implements memmap.File.FD. +func (mf *accelFDMemmapFile) FD() int { + return int(mf.fd.hostFD) +} diff --git a/pkg/sentry/devices/accel/accel_unsafe.go b/pkg/sentry/devices/accel/accel_unsafe.go new file mode 100644 index 000000000..c3e04336d --- /dev/null +++ b/pkg/sentry/devices/accel/accel_unsafe.go @@ -0,0 +1,35 @@ +// 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 accel + +import ( + "unsafe" + + "golang.org/x/exp/constraints" + "golang.org/x/sys/unix" + "gvisor.dev/gvisor/pkg/abi/gasket" +) + +func ioctlInvokePtrArg[Params any](hostFd int32, cmd gasket.Ioctl, params *Params) (uintptr, error) { + return ioctlInvoke[uintptr](hostFd, cmd, uintptr(unsafe.Pointer(params))) +} + +func ioctlInvoke[Arg constraints.Integer](hostFd int32, cmd gasket.Ioctl, arg Arg) (uintptr, error) { + n, _, errno := unix.RawSyscall(unix.SYS_IOCTL, uintptr(hostFd), uintptr(cmd), uintptr(arg)) + if errno != 0 { + return n, errno + } + return n, nil +} diff --git a/pkg/sentry/devices/accel/device.go b/pkg/sentry/devices/accel/device.go index 56fbe44de..56339c853 100644 --- a/pkg/sentry/devices/accel/device.go +++ b/pkg/sentry/devices/accel/device.go @@ -20,18 +20,28 @@ import ( "golang.org/x/sys/unix" "gvisor.dev/gvisor/pkg/abi/linux" "gvisor.dev/gvisor/pkg/context" + "gvisor.dev/gvisor/pkg/fdnotifier" "gvisor.dev/gvisor/pkg/sentry/fsimpl/devtmpfs" "gvisor.dev/gvisor/pkg/sentry/vfs" + "gvisor.dev/gvisor/pkg/sync" ) // accelDevice implements vfs.Device for /dev/accel[0-9]+. // // +stateify savable type accelDevice struct { + mu sync.Mutex + minor uint32 + // +checklocks:mu + openWriteFDs uint32 + // +checklocks:mu + devAddrSet DevAddrSet } func (dev *accelDevice) 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) hostFD, err := unix.Openat(-1, hostPath, int((opts.Flags&unix.O_ACCMODE)|unix.O_NOFOLLOW), 0) if err != nil { @@ -40,6 +50,7 @@ func (dev *accelDevice) Open(ctx context.Context, mnt *vfs.Mount, vfsd *vfs.Dent } fd := &accelFD{ hostFD: int32(hostFD), + device: dev, } if err := fd.vfsfd.Init(fd, opts.Flags, mnt, vfsd, &vfs.FileDescriptionOptions{ UseDentryMetadata: true, @@ -47,6 +58,14 @@ func (dev *accelDevice) Open(ctx context.Context, mnt *vfs.Mount, vfsd *vfs.Dent unix.Close(hostFD) return nil, err } + if err := fdnotifier.AddFD(int32(hostFD), &fd.queue); err != nil { + unix.Close(hostFD) + return nil, err + } + fd.memmapFile.fd = fd + if vfs.MayWriteFileWithOpenFlags(opts.Flags) { + dev.openWriteFDs++ + } return &fd.vfsfd, nil } diff --git a/pkg/sentry/devices/accel/gasket.go b/pkg/sentry/devices/accel/gasket.go new file mode 100644 index 000000000..98e59f1bc --- /dev/null +++ b/pkg/sentry/devices/accel/gasket.go @@ -0,0 +1,167 @@ +// 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 accel + +import ( + "fmt" + + "golang.org/x/sys/unix" + "gvisor.dev/gvisor/pkg/abi/gasket" + "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/hostarch" + "gvisor.dev/gvisor/pkg/sentry/fsimpl/eventfd" + "gvisor.dev/gvisor/pkg/sentry/kernel" + "gvisor.dev/gvisor/pkg/sentry/memmap" + "gvisor.dev/gvisor/pkg/sentry/mm" +) + +func gasketMapBufferIoctl(ctx context.Context, t *kernel.Task, hostFd int32, fd *accelFD, paramsAddr hostarch.Addr) (uintptr, error) { + var userIoctlParams gasket.GasketPageTableIoctl + if _, err := userIoctlParams.CopyIn(t, paramsAddr); err != nil { + return 0, err + } + + tmm := t.MemoryManager() + ar, ok := tmm.CheckIORange(hostarch.Addr(userIoctlParams.HostAddress), int64(userIoctlParams.Size)) + if !ok { + return 0, linuxerr.EFAULT + } + + if !ar.IsPageAligned() || (userIoctlParams.Size/hostarch.PageSize) == 0 { + 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 { + 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 /* ignorePermissions */) + 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{pr.Offset, 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 /* old_size */, uintptr(im.Len()), linux.MREMAP_MAYMOVE|linux.MREMAP_FIXED, sentryAddr, 0); errno != 0 { + return 0, errno + } + sentryAddr += uintptr(im.Len()) + ims = ims.Tail() + } + } + sentryIoctlParams := userIoctlParams + sentryIoctlParams.HostAddress = uint64(m) + n, err := ioctlInvokePtrArg(hostFd, gasket.GASKET_IOCTL_MAP_BUFFER, &sentryIoctlParams) + 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() + devAddr := userIoctlParams.DeviceAddress + for _, pr := range prs { + rlen := uint64(pr.Source.Length()) + if !fd.device.devAddrSet.Add(DevAddrRange{ + devAddr, + devAddr + rlen, + }, pinnedAccelMem{pinnedRange: pr, pageTableIndex: userIoctlParams.PageTableIndex}) { + panic(fmt.Sprintf("unexpected overlap of devaddr range [%#x-%#x)", devAddr, devAddr+rlen)) + } + devAddr += rlen + } + return n, nil +} + +func gasketUnmapBufferIoctl(ctx context.Context, t *kernel.Task, hostFd int32, fd *accelFD, paramsAddr hostarch.Addr) (uintptr, error) { + var userIoctlParams gasket.GasketPageTableIoctl + if _, err := userIoctlParams.CopyIn(t, paramsAddr); err != nil { + return 0, err + } + sentryIoctlParams := userIoctlParams + sentryIoctlParams.HostAddress = 0 // clobber this value, it's unused. + n, err := ioctlInvokePtrArg(hostFd, gasket.GASKET_IOCTL_UNMAP_BUFFER, &sentryIoctlParams) + if err != nil { + return n, err + } + fd.device.mu.Lock() + defer fd.device.mu.Unlock() + s := &fd.device.devAddrSet + r := DevAddrRange{userIoctlParams.DeviceAddress, userIoctlParams.DeviceAddress + userIoctlParams.Size} + seg := s.LowerBoundSegment(r.Start) + for seg.Ok() && seg.Start() < r.End { + seg = s.Isolate(seg, r) + v := seg.Value() + mm.Unpin([]mm.PinnedRange{v.pinnedRange}) + gap := s.Remove(seg) + seg = gap.NextSegment() + } + return n, nil +} + +func gasketInterruptMappingIoctl(ctx context.Context, t *kernel.Task, hostFd int32, paramsAddr hostarch.Addr) (uintptr, error) { + var userIoctlParams gasket.GasketInterruptMapping + if _, err := userIoctlParams.CopyIn(t, paramsAddr); err != nil { + return 0, err + } + + // Check that 'userEventFD.Eventfd' is an eventfd. + eventFileGeneric, _ := t.FDTable().Get(int32(userIoctlParams.EventFD)) + if eventFileGeneric == nil { + return 0, linuxerr.EBADF + } + defer eventFileGeneric.DecRef(ctx) + eventFile, ok := eventFileGeneric.Impl().(*eventfd.EventFileDescription) + if !ok { + return 0, linuxerr.EINVAL + } + + eventfd, err := eventFile.HostFD() + if err != nil { + return 0, err + } + + sentryIoctlParams := userIoctlParams + sentryIoctlParams.EventFD = uint64(eventfd) + n, err := ioctlInvokePtrArg(hostFd, gasket.GASKET_IOCTL_REGISTER_INTERRUPT, &sentryIoctlParams) + if err != nil { + return n, err + } + + outIoctlParams := sentryIoctlParams + outIoctlParams.EventFD = userIoctlParams.EventFD + if _, err := outIoctlParams.CopyOut(t, paramsAddr); err != nil { + return n, err + } + return n, nil +}