mirror of
https://github.com/netbirdio/gvisor.git
synced 2026-05-22 17:12:49 -07:00
nvproxy: implement frontendFDMemmapFile.MapInternal()
New test, before this CL: ``` Testing read/write syscalls on cudaMallocHost memory cuda_malloc: write: Bad address ``` After this CL: ``` Testing read/write syscalls on cudaMallocHost memory Testing cudaMallocManaged(flags=cudaMemAttachGlobal) Testing cudaMallocManaged(flags=cudaMemAttachGlobal) with prefetching Testing cudaMallocManaged(flags=cudaMemAttachHost) Testing cudaMallocManaged(flags=cudaMemAttachHost) with prefetching Testing read/write syscalls on cudaMallocManaged memory All tests passed ``` Fixes #10879 PiperOrigin-RevId: 672721411
This commit is contained in:
@@ -29,6 +29,106 @@ __global__ void addKernel(std::uint32_t* data) {
|
||||
data[index] += static_cast<std::uint32_t>(index);
|
||||
}
|
||||
|
||||
void TestMallocHostReadWrite(int device) {
|
||||
constexpr size_t kNumBlocks = 32;
|
||||
constexpr size_t kNumThreads = 64;
|
||||
constexpr size_t kNumElems = kNumBlocks * kNumThreads;
|
||||
|
||||
constexpr size_t kNumBytes = kNumElems * sizeof(std::uint32_t);
|
||||
std::uint32_t* cpu_data = nullptr;
|
||||
CHECK_CUDA(cudaMallocHost(&cpu_data, kNumBytes, cudaHostAllocWriteCombined));
|
||||
std::uint32_t* gpu_data = nullptr;
|
||||
CHECK_CUDA(cudaMalloc(&gpu_data, kNumBytes));
|
||||
|
||||
// Initialize all elements in the host array with a random value.
|
||||
std::random_device rd;
|
||||
const std::uint32_t init_val =
|
||||
std::uniform_int_distribution<std::uint32_t>()(rd);
|
||||
for (size_t i = 0; i < kNumElems; i++) {
|
||||
cpu_data[i] = init_val;
|
||||
}
|
||||
|
||||
// Write the host array's contents to a temporary file.
|
||||
char filename[] = "/tmp/cudaMallocHostTest.XXXXXX";
|
||||
int fd = mkstemp(filename);
|
||||
if (fd < 0) {
|
||||
err(1, "mkstemp");
|
||||
}
|
||||
size_t done = 0;
|
||||
while (done < kNumBytes) {
|
||||
ssize_t n = write(fd, reinterpret_cast<char*>(cpu_data) + done,
|
||||
kNumBytes - done);
|
||||
if (n >= 0) {
|
||||
done += n;
|
||||
} else if (n < 0 && errno != EINTR) {
|
||||
err(1, "write");
|
||||
}
|
||||
}
|
||||
|
||||
// Copy the array to the device, mutate it there, and copy it back.
|
||||
CHECK_CUDA(cudaMemcpy(gpu_data, cpu_data, kNumBytes, cudaMemcpyHostToDevice));
|
||||
addKernel<<<kNumBlocks, kNumThreads>>>(gpu_data);
|
||||
CHECK_CUDA(cudaDeviceSynchronize());
|
||||
CHECK_CUDA(cudaMemcpy(cpu_data, gpu_data, kNumBytes, cudaMemcpyDeviceToHost));
|
||||
|
||||
// Check that the array has the expected result.
|
||||
for (size_t i = 0; i < kNumElems; i++) {
|
||||
std::uint32_t want = init_val + static_cast<std::uint32_t>(i);
|
||||
if (cpu_data[i] != want) {
|
||||
std::cout << "cpu_data[" << i << "]: got " << cpu_data[i] << ", wanted "
|
||||
<< want << " = " << init_val << " + " << i << std::endl;
|
||||
abort();
|
||||
}
|
||||
}
|
||||
|
||||
// Read the array's original contents back from the temporary file.
|
||||
if (lseek(fd, 0, SEEK_SET) < 0) {
|
||||
err(1, "lseek");
|
||||
}
|
||||
done = 0;
|
||||
while (done < kNumBytes) {
|
||||
ssize_t n = read(fd, reinterpret_cast<char*>(cpu_data) + done,
|
||||
kNumBytes - done);
|
||||
if (n > 0) {
|
||||
done += n;
|
||||
} else if (n == 0) {
|
||||
errx(1, "read: unexpected EOF after %zu bytes", done);
|
||||
} else if (n < 0 && errno != EINTR) {
|
||||
err(1, "read");
|
||||
}
|
||||
}
|
||||
|
||||
// Check that the array matches what we originally wrote.
|
||||
for (size_t i = 0; i < kNumElems; i++) {
|
||||
std::uint32_t want = init_val;
|
||||
if (cpu_data[i] != want) {
|
||||
std::cout << "cpu_data[" << i << "]: got " << cpu_data[i] << ", wanted "
|
||||
<< want << " = " << init_val << " + " << i << std::endl;
|
||||
abort();
|
||||
}
|
||||
}
|
||||
|
||||
// Mutate the array on the device again.
|
||||
CHECK_CUDA(cudaMemcpy(gpu_data, cpu_data, kNumBytes, cudaMemcpyHostToDevice));
|
||||
addKernel<<<kNumBlocks, kNumThreads>>>(gpu_data);
|
||||
CHECK_CUDA(cudaDeviceSynchronize());
|
||||
CHECK_CUDA(cudaMemcpy(cpu_data, gpu_data, kNumBytes, cudaMemcpyDeviceToHost));
|
||||
|
||||
// Check that the array has the expected result again.
|
||||
for (size_t i = 0; i < kNumElems; i++) {
|
||||
std::uint32_t want = init_val + static_cast<std::uint32_t>(i);
|
||||
if (cpu_data[i] != want) {
|
||||
std::cout << "cpu_data[" << i << "]: got " << cpu_data[i] << ", wanted "
|
||||
<< want << " = " << init_val << " + " << i << std::endl;
|
||||
abort();
|
||||
}
|
||||
}
|
||||
|
||||
close(fd);
|
||||
CHECK_CUDA(cudaFreeHost(cpu_data));
|
||||
CHECK_CUDA(cudaFree(gpu_data));
|
||||
}
|
||||
|
||||
void TestMallocManagedRoundTrip(int device, unsigned int malloc_flags,
|
||||
bool prefetch) {
|
||||
constexpr size_t kNumBlocks = 32;
|
||||
@@ -169,6 +269,10 @@ int main() {
|
||||
int device;
|
||||
CHECK_CUDA(cudaGetDevice(&device));
|
||||
|
||||
std::cout << "Testing read/write syscalls on cudaMallocHost memory"
|
||||
<< std::endl;
|
||||
TestMallocHostReadWrite(device);
|
||||
|
||||
std::cout << "Testing cudaMallocManaged(flags=cudaMemAttachGlobal)"
|
||||
<< std::endl;
|
||||
TestMallocManagedRoundTrip(device, cudaMemAttachGlobal, false);
|
||||
@@ -17,5 +17,5 @@
|
||||
set -eux
|
||||
|
||||
cd /
|
||||
nvcc cuda_malloc_managed.cu -o cuda_malloc_managed
|
||||
./cuda_malloc_managed
|
||||
nvcc cuda_malloc.cu -o cuda_malloc
|
||||
./cuda_malloc
|
||||
|
||||
@@ -190,6 +190,7 @@ analyzers:
|
||||
- "pkg/flipcall/.*_unsafe.go" # Special case.
|
||||
- pkg/gohacks/noescape_unsafe.go # Special case.
|
||||
- pkg/ring0/pagetables/allocator_unsafe.go # Special case.
|
||||
- pkg/sentry/devices/nvproxy/frontend_mmap_unsafe.go # Special case.
|
||||
- pkg/sentry/fsutil/host_file_mapper_unsafe.go # Special case.
|
||||
- pkg/sentry/pgalloc/pgalloc_unsafe.go # Special case.
|
||||
- pkg/sentry/platform/kvm/bluepill_unsafe.go # Special case.
|
||||
|
||||
@@ -13,6 +13,13 @@ declare_mutex(
|
||||
prefix = "fds",
|
||||
)
|
||||
|
||||
declare_mutex(
|
||||
name = "frontend_mmap_mutex",
|
||||
out = "frontend_mmap_mutex.go",
|
||||
package = "nvproxy",
|
||||
prefix = "frontendMmap",
|
||||
)
|
||||
|
||||
declare_mutex(
|
||||
name = "objs_mutex",
|
||||
out = "objs_mutex.go",
|
||||
@@ -38,6 +45,8 @@ go_library(
|
||||
"fds_mutex.go",
|
||||
"frontend.go",
|
||||
"frontend_mmap.go",
|
||||
"frontend_mmap_mutex.go",
|
||||
"frontend_mmap_unsafe.go",
|
||||
"frontend_unsafe.go",
|
||||
"nvproxy.go",
|
||||
"nvproxy_unsafe.go",
|
||||
|
||||
@@ -128,7 +128,13 @@ type frontendFD struct {
|
||||
cachedEvents atomicbitops.Uint64
|
||||
appQueue waiter.Queue
|
||||
|
||||
haveMmapContext atomicbitops.Bool `state:"nosave"`
|
||||
// mmapMu protects the following fields.
|
||||
mmapMu frontendMmapMutex `state:"nosave"`
|
||||
// These fields are marked nosave since we do not automatically reinvoke
|
||||
// NV_ESC_RM_MAP_MEMORY after restore, so restored FDs have no
|
||||
// mmap_context.
|
||||
mmapLength uint64 `state:"nosave"`
|
||||
mmapInternal uintptr `state:"nosave"`
|
||||
|
||||
// clients are handles of clients owned by this frontendFD. clients is
|
||||
// protected by dev.nvp.objsMu.
|
||||
@@ -137,6 +143,12 @@ type frontendFD struct {
|
||||
|
||||
// Release implements vfs.FileDescriptionImpl.Release.
|
||||
func (fd *frontendFD) Release(ctx context.Context) {
|
||||
fd.mmapMu.Lock()
|
||||
if fd.mmapInternal != 0 {
|
||||
unix.RawSyscall(unix.SYS_MUNMAP, fd.mmapInternal, uintptr(fd.mmapLength), 0)
|
||||
}
|
||||
fd.mmapMu.Unlock()
|
||||
|
||||
fdnotifier.RemoveFD(fd.hostFD)
|
||||
fd.appQueue.Notify(waiter.EventHUp)
|
||||
|
||||
@@ -1020,17 +1032,21 @@ func rmMapMemory(fi *frontendIoctlState) (uintptr, error) {
|
||||
if !ok {
|
||||
return 0, linuxerr.EINVAL
|
||||
}
|
||||
if mapFile.haveMmapContext.Load() || !mapFile.haveMmapContext.CompareAndSwap(false, true) {
|
||||
|
||||
mapFile.mmapMu.Lock()
|
||||
defer mapFile.mmapMu.Unlock()
|
||||
if mapFile.mmapLength != 0 {
|
||||
fi.ctx.Warningf("nvproxy: attempted to reuse FD %d for NV_ESC_RM_MAP_MEMORY", ioctlParams.FD)
|
||||
return 0, linuxerr.EINVAL
|
||||
}
|
||||
|
||||
origFD := ioctlParams.FD
|
||||
ioctlParams.FD = mapFile.hostFD
|
||||
|
||||
n, err := frontendIoctlInvoke(fi, &ioctlParams)
|
||||
if err != nil {
|
||||
return n, err
|
||||
}
|
||||
mapFile.mmapLength = ioctlParams.Params.Length
|
||||
|
||||
ioctlParams.FD = origFD
|
||||
if _, err := ioctlParams.CopyOut(fi.t, fi.ioctlParamsAddr); err != nil {
|
||||
|
||||
@@ -16,10 +16,7 @@ package nvproxy
|
||||
|
||||
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"
|
||||
)
|
||||
@@ -75,13 +72,6 @@ func (mf *frontendFDMemmapFile) IncRef(fr memmap.FileRange, memCgID uint32) {
|
||||
func (mf *frontendFDMemmapFile) DecRef(fr memmap.FileRange) {
|
||||
}
|
||||
|
||||
// MapInternal implements memmap.File.MapInternal.
|
||||
func (mf *frontendFDMemmapFile) MapInternal(fr memmap.FileRange, at hostarch.AccessType) (safemem.BlockSeq, error) {
|
||||
// FIXME(jamieliu): determine if this is safe
|
||||
log.Traceback("nvproxy: rejecting frontendFDMemmapFile.MapInternal")
|
||||
return safemem.BlockSeq{}, linuxerr.EINVAL
|
||||
}
|
||||
|
||||
// FD implements memmap.File.FD.
|
||||
func (mf *frontendFDMemmapFile) FD() int {
|
||||
return int(mf.fd.hostFD)
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
// Copyright 2024 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 nvproxy
|
||||
|
||||
import (
|
||||
"unsafe"
|
||||
|
||||
"golang.org/x/sys/unix"
|
||||
"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"
|
||||
)
|
||||
|
||||
// MapInternal implements memmap.File.MapInternal.
|
||||
func (mf *frontendFDMemmapFile) MapInternal(fr memmap.FileRange, at hostarch.AccessType) (safemem.BlockSeq, error) {
|
||||
if at.Execute {
|
||||
return safemem.BlockSeq{}, linuxerr.EACCES
|
||||
}
|
||||
|
||||
mf.fd.mmapMu.Lock()
|
||||
defer mf.fd.mmapMu.Unlock()
|
||||
if mf.fd.mmapInternal == 0 {
|
||||
if mf.fd.mmapLength == 0 {
|
||||
// This shouldn't be possible.
|
||||
log.Traceback("nvproxy: frontendFDMemmapFile.MapInternal() called before NV_ESC_RM_MAP_MEMORY")
|
||||
return safemem.BlockSeq{}, linuxerr.EINVAL
|
||||
}
|
||||
// Nvidia kernel driver:
|
||||
// kernel-open/nvidia/nv-mmap.c:nvidia_mmap_helper() requires vm_pgoff
|
||||
// == 0 (so we must pass offset 0 here), and conditionally requires
|
||||
// NV_VMA_SIZE(vma) == mmap_context->mmap_size (so we pass length
|
||||
// mmapLength here).
|
||||
m, _, errno := unix.Syscall6(unix.SYS_MMAP, 0 /* addr */, uintptr(mf.fd.mmapLength), unix.PROT_READ|unix.PROT_WRITE, unix.MAP_SHARED, uintptr(mf.fd.hostFD), 0 /* offset */)
|
||||
if errno != 0 {
|
||||
return safemem.BlockSeq{}, errno
|
||||
}
|
||||
mf.fd.mmapInternal = m
|
||||
}
|
||||
mappedFR := memmap.FileRange{0, mf.fd.mmapLength}
|
||||
if !mappedFR.IsSupersetOf(fr) {
|
||||
return safemem.BlockSeq{}, linuxerr.EINVAL
|
||||
}
|
||||
// mmap_context::prot is determined internally during NV_ESC_RM_MAP_MEMORY
|
||||
// (see
|
||||
// src/nvidia/arch/nvalloc/unix/src/osapi.c:RmCreateMmapContextLocked());
|
||||
// nvidia_mmap_helper() propagates this to vm_area_struct::vm_page_prot, so
|
||||
// PROT_WRITE on a read-only mapping will succeed at mmap time but fault at
|
||||
// write time. Thus, these mappings should use safecopy (i.e.
|
||||
// BlockFromUnsafePointer rather than BlockFromSafePointer).
|
||||
return safemem.BlockSeqOf(safemem.BlockFromUnsafePointer(unsafe.Pointer(mf.fd.mmapInternal+uintptr(fr.Start)), int(fr.Length()))), nil
|
||||
}
|
||||
Reference in New Issue
Block a user