platform/systrap: return memory access type based on a page fault error code

Now we don't need to trigger a second fault to figure out whether it was write
or read access.

Fixes #11008

Co-developed-by: Jamie Liu <jamieliu@google.com>
PiperOrigin-RevId: 697677262
This commit is contained in:
Andrei Vagin
2024-11-18 10:33:59 -08:00
committed by gVisor bot
parent 336dc85043
commit 03a28d158e
11 changed files with 106 additions and 130 deletions
+9
View File
@@ -22,3 +22,12 @@ package linux
//
// The array has to be sorted in decreasing order.
var feasibleTaskSizes = []uintptr{0xfffffffffff000, 0x7ffffffff000}
// Page fault error codes
const (
X86_PF_PROT = 1 << iota
X86_PF_WRITE
X86_PF_USER
X86_PF_RSVD
X86_PF_INSTR
)
+33
View File
@@ -53,6 +53,39 @@ var (
ByteOrder = binary.LittleEndian
)
// Arm64: Exception Syndrome Register EL1.
const (
_ESR_ELx_EC_SHIFT = 26
_ESR_ELx_EC_MASK = 0x3F << _ESR_ELx_EC_SHIFT
_ESR_ELx_EC_IABT_LOW = 0x20
_ESR_ELx_EC_DABT_LOW = 0x24
_ESR_ELx_WNR = 1 << 6
_ESR_ELx_CM = 1 << 8
)
// ESRAccessType returns the memory access type for the given ESR (Exception
// Syndrome Register) code. If code does not represent an invalid memory
// access from a lower exception level, ESRAccessType returns NoAccess.
//
//go:nosplit
func ESRAccessType(code uint64) AccessType {
switch (code & _ESR_ELx_EC_MASK) >> _ESR_ELx_EC_SHIFT {
case _ESR_ELx_EC_IABT_LOW:
return Execute
case _ESR_ELx_EC_DABT_LOW:
// For faults on cache maintenance and address translation
// instructions, _ESR_ELx_WNR is always set.
if code&(_ESR_ELx_WNR|_ESR_ELx_CM) == _ESR_ELx_WNR {
return Write
}
return Read
default:
return NoAccess
}
}
func init() {
// Make sure the page size is 4K on arm64 platform.
if size := unix.Getpagesize(); size != PageSize {
+1 -24
View File
@@ -148,25 +148,6 @@ func nonCanonical(addr uint64, signal int32, info *linux.SignalInfo) (hostarch.A
return hostarch.NoAccess, platform.ErrContextSignal
}
// isInstructionAbort returns true if it is an instruction abort.
//
//go:nosplit
func isInstructionAbort(code uint64) bool {
value := (code & _ESR_ELx_EC_MASK) >> _ESR_ELx_EC_SHIFT
return value == _ESR_ELx_EC_IABT_LOW
}
// isWriteFault returns whether it is a write fault.
//
//go:nosplit
func isWriteFault(code uint64) bool {
if isInstructionAbort(code) {
return false
}
return (code & _ESR_ELx_WNR) != 0
}
// fault generates an appropriate fault return.
//
//go:nosplit
@@ -186,11 +167,7 @@ func (c *vCPU) fault(signal int32, info *linux.SignalInfo) (hostarch.AccessType,
info.SetAddr(uint64(faultAddr))
accessType := hostarch.AccessType{}
if signal == int32(unix.SIGSEGV) {
accessType = hostarch.AccessType{
Read: !isWriteFault(uint64(code)),
Write: isWriteFault(uint64(code)),
Execute: isInstructionAbort(uint64(code)),
}
accessType = hostarch.ESRAccessType(uint64(code))
}
ret := code & _ESR_ELx_FSC
+14 -28
View File
@@ -782,7 +782,7 @@ func (s *subprocess) decAwakeContexts() {
// This function returns true on a system call, false on a signal.
// The second return value is true if a syscall instruction can be replaced on
// a function call.
func (s *subprocess) switchToApp(c *platformContext, ac *arch.Context64) (isSyscall bool, shouldPatchSyscall bool, err *platform.ContextError) {
func (s *subprocess) switchToApp(c *platformContext, ac *arch.Context64) (isSyscall bool, shouldPatchSyscall bool, at hostarch.AccessType, err *platform.ContextError) {
// Reset necessary registers.
regs := &ac.StateData().Regs
s.resetSysemuRegs(regs)
@@ -795,7 +795,7 @@ func (s *subprocess) switchToApp(c *platformContext, ac *arch.Context64) (isSysc
// Pending interrupt; simulate.
ctx.clearInterrupt()
c.signalInfo = linux.SignalInfo{Signo: int32(platform.SignalInterrupt)}
return false, false, nil
return false, false, hostarch.NoAccess, nil
}
defer func() {
ctx.clearInterrupt()
@@ -811,20 +811,20 @@ func (s *subprocess) switchToApp(c *platformContext, ac *arch.Context64) (isSysc
}
ctx.setState(sysmsg.ContextStateNone)
if err := s.contextQueue.add(ctx); err != nil {
return false, false, err
return false, false, hostarch.NoAccess, err
}
if err := s.waitOnState(ctx); err != nil {
return false, false, corruptedSharedMemoryErr(err.Error())
return false, false, hostarch.NoAccess, corruptedSharedMemoryErr(err.Error())
}
// Check if there's been an error.
threadID := ctx.threadID()
if threadID != invalidThreadID {
if sysThread, ok := s.sysmsgThreads[threadID]; ok && sysThread.msg.Err != 0 {
return false, false, sysThread.msg.ConvertSysmsgErr()
return false, false, hostarch.NoAccess, sysThread.msg.ConvertSysmsgErr()
}
return false, false, corruptedSharedMemoryErr(fmt.Sprintf("found unexpected ThreadContext.ThreadID field, expected %d found %d", invalidThreadID, threadID))
return false, false, hostarch.NoAccess, corruptedSharedMemoryErr(fmt.Sprintf("found unexpected ThreadContext.ThreadID field, expected %d found %d", invalidThreadID, threadID))
}
// Copy register state locally.
@@ -843,15 +843,19 @@ func (s *subprocess) switchToApp(c *platformContext, ac *arch.Context64) (isSysc
if ctxState == sysmsg.ContextStateSyscall || ctxState == sysmsg.ContextStateSyscallTrap {
if maybePatchSignalInfo(regs, &c.signalInfo) {
return false, false, nil
return false, false, hostarch.Execute, nil
}
updateSyscallRegs(regs)
return true, shouldPatchSyscall, nil
return true, shouldPatchSyscall, hostarch.NoAccess, nil
} else if ctxState != sysmsg.ContextStateFault {
return false, false, corruptedSharedMemoryErr(fmt.Sprintf("unknown context state: %v", ctxState))
return false, false, hostarch.NoAccess, corruptedSharedMemoryErr(fmt.Sprintf("unknown context state: %v", ctxState))
}
return false, false, nil
at = hostarch.NoAccess
if c.signalInfo.Signo == int32(linux.SIGSEGV) {
at = sigErrorToAccessType(ctx.shared.SigError)
}
return false, false, at, nil
}
func (s *subprocess) waitOnState(ctx *sharedContext) error {
@@ -989,24 +993,6 @@ func (s *subprocess) MapFile(addr hostarch.Addr, f memmap.File, fr memmap.FileRa
// Unmap implements platform.AddressSpace.Unmap.
func (s *subprocess) Unmap(addr hostarch.Addr, length uint64) {
ar, ok := addr.ToRange(length)
if !ok {
panic(fmt.Sprintf("addr %#x + length %#x overflows", addr, length))
}
s.mu.Lock()
for c := range s.faultedContexts {
c.mu.Lock()
if c.lastFaultSP == s && ar.Contains(c.lastFaultAddr) {
// Forget the last fault so that if c faults again, the fault isn't
// incorrectly reported as a write fault. If this is being called
// due to munmap() of the corresponding vma, handling of the second
// fault will fail anyway.
c.lastFaultSP = nil
delete(s.faultedContexts, c)
}
c.mu.Unlock()
}
s.mu.Unlock()
_, err := s.syscall(
unix.SYS_MUNMAP,
arch.SyscallArgument{Value: uintptr(addr)},
@@ -23,6 +23,7 @@ import (
"golang.org/x/sys/unix"
"gvisor.dev/gvisor/pkg/abi/linux"
"gvisor.dev/gvisor/pkg/hostarch"
"gvisor.dev/gvisor/pkg/hostsyscall"
"gvisor.dev/gvisor/pkg/seccomp"
"gvisor.dev/gvisor/pkg/sentry/arch"
@@ -223,3 +224,16 @@ func retrieveArchSpecificState(ctx *sysmsg.ThreadContext, ac *arch.Context64) {
func archSpecificSysmsgThreadInit(sysThread *sysmsgThread) {
}
func sigErrorToAccessType(sigError uint64) hostarch.AccessType {
switch {
case sigError&linux.X86_PF_WRITE != 0:
return hostarch.Write
case sigError&linux.X86_PF_INSTR != 0:
return hostarch.Execute
case sigError&linux.X86_PF_USER != 0:
return hostarch.Read
default:
return hostarch.NoAccess
}
}
@@ -23,6 +23,7 @@ import (
"golang.org/x/sys/unix"
"gvisor.dev/gvisor/pkg/abi/linux"
"gvisor.dev/gvisor/pkg/hostarch"
"gvisor.dev/gvisor/pkg/hostsyscall"
"gvisor.dev/gvisor/pkg/seccomp"
"gvisor.dev/gvisor/pkg/sentry/arch"
@@ -195,3 +196,7 @@ func archSpecificSysmsgThreadInit(sysThread *sysmsgThread) {
panic(fmt.Sprintf("tkill failed: %v", e))
}
}
func sigErrorToAccessType(sigError uint64) hostarch.AccessType {
return hostarch.ESRAccessType(sigError)
}
@@ -32,6 +32,10 @@
#include "sysmsg_offsets.h"
#include "sysmsg_offsets_amd64.h"
#ifndef X86_TRAP_PF
#define X86_TRAP_PF 14
#endif
// TODO(b/271631387): These globals are shared between AMD64 and ARM64; move to
// sysmsg_lib.c.
struct arch_state __export_arch_state;
@@ -230,6 +234,7 @@ void __export_sighandler(int signo, siginfo_t *siginfo, void *_ucontext) {
if (signo != SIGCHLD ||
ucontext->uc_mcontext.gregs[REG_RIP] < __export_stub_start) {
ctx->ptregs.fs_base = fs_base;
ctx->err = 0;
gregs_to_ptregs(ucontext, &ctx->ptregs);
memcpy(ctx->fpstate, (uint8_t *)ucontext->uc_mcontext.fpregs,
__export_arch_state.fp_len);
@@ -307,8 +312,12 @@ void __export_sighandler(int signo, siginfo_t *siginfo, void *_ucontext) {
ctx->ptregs.orig_rax += 0x86000000;
break;
}
case SIGCHLD:
case SIGSEGV:
if (ucontext->uc_mcontext.gregs[REG_TRAPNO] == X86_TRAP_PF) {
ctx->err = ucontext->uc_mcontext.gregs[REG_ERR];
}
// fallthrough
case SIGCHLD:
case SIGBUS:
case SIGFPE:
case SIGTRAP:
@@ -139,6 +139,7 @@ void __export_sighandler(int signo, siginfo_t *siginfo, void *_ucontext) {
memcpy(ctx->fpstate, fpStatePointer, kFpsimdContextSize);
ctx->tls = get_tls();
ctx->siginfo = *siginfo;
ctx->err = 0;
switch (signo) {
case SIGSYS: {
ctx_state = CONTEXT_STATE_SYSCALL;
@@ -151,9 +152,22 @@ void __export_sighandler(int signo, siginfo_t *siginfo, void *_ucontext) {
}
break;
}
case SIGCHLD:
case SIGSEGV:
case SIGSEGV: {
unsigned char *base = &ucontext->uc_mcontext.__reserved[0];
size_t offset = 0;
while (1) {
struct _aarch64_ctx *head = (struct _aarch64_ctx *)(base + offset);
if (head->magic == ESR_MAGIC) {
ctx->err = ((struct esr_context *)head)->esr;
break;
}
if (head->magic == 0 || head->magic == EXTRA_MAGIC) break;
offset += head->size;
}
}
// fallthrough
case SIGBUS:
case SIGCHLD:
case SIGFPE:
case SIGTRAP:
case SIGILL:
@@ -269,6 +269,8 @@ type ThreadContext struct {
TLS uint64
// Debug is a variable to use to get visibility into the stub from the sentry.
Debug uint64
// SigError is an error code that clarifies the nature of the signal.
SigError uint64
}
// StubError are values that represent known stub-thread failure modes.
@@ -98,6 +98,7 @@ struct thread_context {
uint64_t state_changed_time;
uint64_t tls;
uint64_t debug;
uint64_t err;
};
enum stub_error {
+1 -75
View File
@@ -129,22 +129,6 @@ type platformContext struct {
// mutex protected.
sharedContext *sharedContext
// mu protects the following fields.
mu sync.Mutex
// If lastFaultSP is non-nil, the last platformContext switch was due to a fault
// received while executing lastFaultSP. Only platformContext.Switch may set
// lastFaultSP to a non-nil value.
lastFaultSP *subprocess
// lastFaultAddr is the last faulting address; this is only meaningful if
// lastFaultSP is non-nil.
lastFaultAddr hostarch.Addr
// lastFaultIP is the address of the last faulting instruction;
// this is also only meaningful if lastFaultSP is non-nil.
lastFaultIP hostarch.Addr
// needRestoreFPState indicates that the FPU state has been changed by
// the Sentry and has to be updated on the stub thread.
needRestoreFPState bool
@@ -182,7 +166,7 @@ func (c *platformContext) Switch(ctx pkgcontext.Context, mm platform.MemoryManag
}
restart:
isSyscall, needPatch, err := s.switchToApp(c, ac)
isSyscall, needPatch, at, err := s.switchToApp(c, ac)
if err != nil {
return nil, hostarch.NoAccess, err
}
@@ -199,76 +183,18 @@ restart:
ctx.Warningf("usertrap.HandleFault failed: %v", err)
}
}
var (
faultSP *subprocess
faultAddr hostarch.Addr
faultIP hostarch.Addr
)
if !isSyscall && linux.Signal(c.signalInfo.Signo) == linux.SIGSEGV {
faultSP = s
faultAddr = hostarch.Addr(c.signalInfo.Addr())
faultIP = hostarch.Addr(ac.IP())
}
// Update the platformContext to reflect the outcome of this context switch.
c.mu.Lock()
lastFaultSP := c.lastFaultSP
lastFaultAddr := c.lastFaultAddr
lastFaultIP := c.lastFaultIP
// At this point, c may not yet be in s.faultedContexts, so c.lastFaultSP won't
// be updated by s.Unmap(). This is fine; we only need to synchronize with
// calls to s.Unmap() that occur after the handling of this fault.
c.lastFaultSP = faultSP
c.lastFaultAddr = faultAddr
c.lastFaultIP = faultIP
c.mu.Unlock()
// Update subprocesses to reflect the outcome of this context switch.
if lastFaultSP != faultSP {
if lastFaultSP != nil {
lastFaultSP.mu.Lock()
delete(lastFaultSP.faultedContexts, c)
lastFaultSP.mu.Unlock()
}
if faultSP != nil {
faultSP.mu.Lock()
faultSP.faultedContexts[c] = struct{}{}
faultSP.mu.Unlock()
}
}
if isSyscall {
return nil, hostarch.NoAccess, nil
}
si := c.signalInfo
if faultSP == nil {
// Non-fault signal.
return &si, hostarch.NoAccess, platform.ErrContextSignal
}
// See if this can be handled as a CPUID exception.
if linux.Signal(si.Signo) == linux.SIGSEGV && platform.TryCPUIDEmulate(ctx, mm, ac) {
goto restart
}
// Got a page fault. Ideally, we'd get real fault type here, but ptrace
// doesn't expose this information. Instead, we use a simple heuristic:
//
// It was an instruction fault iff the faulting addr == instruction
// pointer.
//
// It was a write fault if the fault is immediately repeated.
at := hostarch.Read
if faultAddr == faultIP {
at.Execute = true
}
if lastFaultSP == faultSP &&
lastFaultAddr == faultAddr &&
lastFaultIP == faultIP {
at.Write = true
}
// Handle as a signal.
return &si, at, platform.ErrContextSignal
}