From 7726eb75058827119e95a6f7c18a9c8f7da51f29 Mon Sep 17 00:00:00 2001 From: Ilya Rezvov Date: Sat, 21 Apr 2018 10:11:37 -0700 Subject: [PATCH] protect memory access with traps --- src/Language/Wasm/Interpreter.hs | 203 ++++++++++++++++++++++--------- tests/Test.hs | 2 +- 2 files changed, 145 insertions(+), 60 deletions(-) diff --git a/src/Language/Wasm/Interpreter.hs b/src/Language/Wasm/Interpreter.hs index c864a22..3270e60 100644 --- a/src/Language/Wasm/Interpreter.hs +++ b/src/Language/Wasm/Interpreter.hs @@ -720,8 +720,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let readByte idx = do byte <- IOVector.read memory $ addr + idx return $ fromIntegral byte `shiftL` (idx * 8) - val <- sum <$> mapM readByte [0..3] - return $ Done ctx { stack = VI32 val : rest } + if addr + 4 > IOVector.length memory + then return Trap + else do + val <- sum <$> mapM readByte [0..3] + return $ Done ctx { stack = VI32 val : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } (I64Load MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -729,8 +732,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let readByte idx = do byte <- IOVector.read memory $ addr + idx return $ fromIntegral byte `shiftL` (idx * 8) - val <- sum <$> mapM readByte [0..7] - return $ Done ctx { stack = VI64 val : rest } + if addr + 8 > IOVector.length memory + then return Trap + else do + val <- sum <$> mapM readByte [0..7] + return $ Done ctx { stack = VI64 val : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } (F32Load MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -738,8 +744,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let readByte idx = do byte <- IOVector.read memory $ addr + idx return $ fromIntegral byte `shiftL` (idx * 8) - val <- wordToFloat . sum <$> mapM readByte [0..3] - return $ Done ctx { stack = VF32 val : rest } + if addr + 4 > IOVector.length memory + then return Trap + else do + val <- wordToFloat . sum <$> mapM readByte [0..3] + return $ Done ctx { stack = VF32 val : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } (F64Load MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -747,21 +756,30 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let readByte idx = do byte <- IOVector.read memory $ addr + idx return $ fromIntegral byte `shiftL` (idx * 8) - val <- wordToDouble . sum <$> mapM readByte [0..7] - return $ Done ctx { stack = VF64 val : rest } + if addr + 8 > IOVector.length memory + then return Trap + else do + val <- wordToDouble . sum <$> mapM readByte [0..7] + return $ Done ctx { stack = VF64 val : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } (I32Load8U MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef let addr = fromIntegral $ v + fromIntegral offset - byte <- IOVector.read memory addr - return $ Done ctx { stack = VI32 (fromIntegral byte) : rest } + if addr + 1 > IOVector.length memory + then return Trap + else do + byte <- IOVector.read memory addr + return $ Done ctx { stack = VI32 (fromIntegral byte) : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } (I32Load8S MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef let addr = fromIntegral $ v + fromIntegral offset - byte <- IOVector.read memory addr - let val = asWord32 $ if byte >= 128 then -1 * fromIntegral (0xFF - byte + 1) else fromIntegral byte - return $ Done ctx { stack = VI32 val : rest } + if addr + 4 > IOVector.length memory + then return Trap + else do + byte <- IOVector.read memory addr + let val = asWord32 $ if byte >= 128 then -1 * fromIntegral (0xFF - byte + 1) else fromIntegral byte + return $ Done ctx { stack = VI32 val : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } (I32Load16U MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -769,8 +787,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let readByte idx = do byte <- IOVector.read memory $ addr + idx return $ fromIntegral byte `shiftL` (idx * 8) - val <- sum <$> mapM readByte [0..1] - return $ Done ctx { stack = VI32 val : rest } + if addr + 2 > IOVector.length memory + then return Trap + else do + val <- sum <$> mapM readByte [0..1] + return $ Done ctx { stack = VI32 val : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } (I32Load16S MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -778,22 +799,31 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let readByte idx = do byte <- IOVector.read memory $ addr + idx return $ (fromIntegral byte :: Word32) `shiftL` (idx * 8) - val <- sum <$> mapM readByte [0..1] - let signed = asWord32 $ if val >= 2 ^ 15 then -1 * fromIntegral (0xFFFF - val + 1) else fromIntegral val - return $ Done ctx { stack = VI32 signed : rest } + if addr + 2 > IOVector.length memory + then return Trap + else do + val <- sum <$> mapM readByte [0..1] + let signed = asWord32 $ if val >= 2 ^ 15 then -1 * fromIntegral (0xFFFF - val + 1) else fromIntegral val + return $ Done ctx { stack = VI32 signed : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } (I64Load8U MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef let addr = fromIntegral $ v + fromIntegral offset - byte <- IOVector.read memory addr - return $ Done ctx { stack = VI64 (fromIntegral byte) : rest } + if addr + 1 > IOVector.length memory + then return Trap + else do + byte <- IOVector.read memory addr + return $ Done ctx { stack = VI64 (fromIntegral byte) : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } (I64Load8S MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef let addr = fromIntegral $ v + fromIntegral offset - byte <- IOVector.read memory addr - let val = asWord64 $ if byte >= 128 then -1 * fromIntegral (0xFF - byte + 1) else fromIntegral byte - return $ Done ctx { stack = VI64 val : rest } + if addr + 1 > IOVector.length memory + then return Trap + else do + byte <- IOVector.read memory addr + let val = asWord64 $ if byte >= 128 then -1 * fromIntegral (0xFF - byte + 1) else fromIntegral byte + return $ Done ctx { stack = VI64 val : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } (I64Load16U MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -801,8 +831,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let readByte idx = do byte <- IOVector.read memory $ addr + idx return $ fromIntegral byte `shiftL` (idx * 8) - val <- sum <$> mapM readByte [0..1] - return $ Done ctx { stack = VI64 val : rest } + if addr + 2 > IOVector.length memory + then return Trap + else do + val <- sum <$> mapM readByte [0..1] + return $ Done ctx { stack = VI64 val : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } (I64Load16S MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -810,9 +843,12 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let readByte idx = do byte <- IOVector.read memory $ addr + idx return $ (fromIntegral byte :: Word32) `shiftL` (idx * 8) - val <- sum <$> mapM readByte [0..1] - let signed = asWord64 $ if val >= 2 ^ 15 then -1 * fromIntegral (0xFFFF - val + 1) else fromIntegral val - return $ Done ctx { stack = VI64 signed : rest } + if addr + 2 > IOVector.length memory + then return Trap + else do + val <- sum <$> mapM readByte [0..1] + let signed = asWord64 $ if val >= 2 ^ 15 then -1 * fromIntegral (0xFFFF - val + 1) else fromIntegral val + return $ Done ctx { stack = VI64 signed : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } (I64Load32U MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -820,8 +856,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let readByte idx = do byte <- IOVector.read memory $ addr + idx return $ fromIntegral byte `shiftL` (idx * 8) - val <- sum <$> mapM readByte [0..3] - return $ Done ctx { stack = VI64 val : rest } + if addr + 4 > IOVector.length memory + then return Trap + else do + val <- sum <$> mapM readByte [0..3] + return $ Done ctx { stack = VI64 val : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } (I64Load32S MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -829,9 +868,12 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let readByte idx = do byte <- IOVector.read memory $ addr + idx return $ (fromIntegral byte :: Word32) `shiftL` (idx * 8) - val <- sum <$> mapM readByte [0..3] - let signed = asWord64 $ fromIntegral $ asInt32 val - return $ Done ctx { stack = VI64 signed : rest } + if addr + 4 > IOVector.length memory + then return Trap + else do + val <- sum <$> mapM readByte [0..3] + let signed = asWord64 $ fromIntegral $ asInt32 val + return $ Done ctx { stack = VI64 signed : rest } step ctx@EvalCtx{ stack = (VI32 v:VI32 va:rest) } (I32Store MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -839,8 +881,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let writeByte idx = do let byte = fromIntegral $ v `shiftR` (idx * 8) .&. 0xFF IOVector.write memory (addr + idx) byte - mapM_ writeByte [0..3] - return $ Done ctx { stack = rest } + if addr + 4 > IOVector.length memory + then return Trap + else do + mapM_ writeByte [0..3] + return $ Done ctx { stack = rest } step ctx@EvalCtx{ stack = (VI64 v:VI32 va:rest) } (I64Store MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -848,8 +893,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let writeByte idx = do let byte = fromIntegral $ v `shiftR` (idx * 8) .&. 0xFF IOVector.write memory (addr + idx) byte - mapM_ writeByte [0..7] - return $ Done ctx { stack = rest } + if addr + 8 > IOVector.length memory + then return Trap + else do + mapM_ writeByte [0..7] + return $ Done ctx { stack = rest } step ctx@EvalCtx{ stack = (VF32 f:VI32 va:rest) } (F32Store MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -858,8 +906,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let writeByte idx = do let byte = fromIntegral $ v `shiftR` (idx * 8) .&. 0xFF IOVector.write memory (addr + idx) byte - mapM_ writeByte [0..3] - return $ Done ctx { stack = rest } + if addr + 4 > IOVector.length memory + then return Trap + else do + mapM_ writeByte [0..3] + return $ Done ctx { stack = rest } step ctx@EvalCtx{ stack = (VF64 f:VI32 va:rest) } (F64Store MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -868,8 +919,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let writeByte idx = do let byte = fromIntegral $ v `shiftR` (idx * 8) .&. 0xFF IOVector.write memory (addr + idx) byte - mapM_ writeByte [0..7] - return $ Done ctx { stack = rest } + if addr + 8 > IOVector.length memory + then return Trap + else do + mapM_ writeByte [0..7] + return $ Done ctx { stack = rest } step ctx@EvalCtx{ stack = (VI32 v:VI32 va:rest) } (I32Store8 MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -877,8 +931,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let writeByte idx = do let byte = fromIntegral $ v `shiftR` (idx * 8) .&. 0xFF IOVector.write memory (addr + idx) byte - mapM_ writeByte [0] - return $ Done ctx { stack = rest } + if addr + 1 > IOVector.length memory + then return Trap + else do + mapM_ writeByte [0] + return $ Done ctx { stack = rest } step ctx@EvalCtx{ stack = (VI32 v:VI32 va:rest) } (I32Store16 MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -886,8 +943,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let writeByte idx = do let byte = fromIntegral $ v `shiftR` (idx * 8) .&. 0xFF IOVector.write memory (addr + idx) byte - mapM_ writeByte [0, 1] - return $ Done ctx { stack = rest } + if addr + 2 > IOVector.length memory + then return Trap + else do + mapM_ writeByte [0, 1] + return $ Done ctx { stack = rest } step ctx@EvalCtx{ stack = (VI64 v:VI32 va:rest) } (I64Store8 MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -895,8 +955,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let writeByte idx = do let byte = fromIntegral $ v `shiftR` (idx * 8) .&. 0xFF IOVector.write memory (addr + idx) byte - mapM_ writeByte [0] - return $ Done ctx { stack = rest } + if addr + 1 > IOVector.length memory + then return Trap + else do + mapM_ writeByte [0] + return $ Done ctx { stack = rest } step ctx@EvalCtx{ stack = (VI64 v:VI32 va:rest) } (I64Store16 MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -904,8 +967,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let writeByte idx = do let byte = fromIntegral $ v `shiftR` (idx * 8) .&. 0xFF IOVector.write memory (addr + idx) byte - mapM_ writeByte [0, 1] - return $ Done ctx { stack = rest } + if addr + 2 > IOVector.length memory + then return Trap + else do + mapM_ writeByte [0, 1] + return $ Done ctx { stack = rest } step ctx@EvalCtx{ stack = (VI64 v:VI32 va:rest) } (I64Store32 MemArg { offset }) = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -913,8 +979,11 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT let writeByte idx = do let byte = fromIntegral $ v `shiftR` (idx * 8) .&. 0xFF IOVector.write memory (addr + idx) byte - mapM_ writeByte [0..3] - return $ Done ctx { stack = rest } + if addr + 4 > IOVector.length memory + then return Trap + else do + mapM_ writeByte [0..3] + return $ Done ctx { stack = rest } step ctx@EvalCtx{ stack = st } CurrentMemory = do let MemoryInstance { memory = memoryRef } = memInstances store ! (memaddrs moduleInstance ! 0) memory <- readIORef memoryRef @@ -1153,21 +1222,37 @@ eval store FunctionInstance { funcType, moduleInstance, code = Function { localT step ctx@EvalCtx{ stack = (VI64 v:rest) } I32WrapI64 = return $ Done ctx { stack = VI32 (fromIntegral $ v .&. 0xFFFFFFFF) : rest } step ctx@EvalCtx{ stack = (VF32 v:rest) } (ITruncFU BS32 BS32) = - return $ Done ctx { stack = VI32 (truncate v) : rest } + if isNaN v + then return Trap + else return $ Done ctx { stack = VI32 (truncate v) : rest } step ctx@EvalCtx{ stack = (VF64 v:rest) } (ITruncFU BS32 BS64) = - return $ Done ctx { stack = VI32 (truncate v) : rest } + if isNaN v + then return Trap + else return $ Done ctx { stack = VI32 (truncate v) : rest } step ctx@EvalCtx{ stack = (VF32 v:rest) } (ITruncFU BS64 BS32) = - return $ Done ctx { stack = VI64 (truncate v) : rest } + if isNaN v + then return Trap + else return $ Done ctx { stack = VI64 (truncate v) : rest } step ctx@EvalCtx{ stack = (VF64 v:rest) } (ITruncFU BS64 BS64) = - return $ Done ctx { stack = VI64 (truncate v) : rest } + if isNaN v + then return Trap + else return $ Done ctx { stack = VI64 (truncate v) : rest } step ctx@EvalCtx{ stack = (VF32 v:rest) } (ITruncFS BS32 BS32) = - return $ Done ctx { stack = VI32 (asWord32 $ truncate v) : rest } + if isNaN v + then return Trap + else return $ Done ctx { stack = VI32 (asWord32 $ truncate v) : rest } step ctx@EvalCtx{ stack = (VF64 v:rest) } (ITruncFS BS32 BS64) = - return $ Done ctx { stack = VI32 (asWord32 $ truncate v) : rest } + if isNaN v + then return Trap + else return $ Done ctx { stack = VI32 (asWord32 $ truncate v) : rest } step ctx@EvalCtx{ stack = (VF32 v:rest) } (ITruncFS BS64 BS32) = - return $ Done ctx { stack = VI64 (asWord64 $ truncate v) : rest } + if isNaN v + then return Trap + else return $ Done ctx { stack = VI64 (asWord64 $ truncate v) : rest } step ctx@EvalCtx{ stack = (VF64 v:rest) } (ITruncFS BS64 BS64) = - return $ Done ctx { stack = VI64 (asWord64 $ truncate v) : rest } + if isNaN v + then return Trap + else return $ Done ctx { stack = VI64 (asWord64 $ truncate v) : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } I64ExtendUI32 = return $ Done ctx { stack = VI64 (fromIntegral v) : rest } step ctx@EvalCtx{ stack = (VI32 v:rest) } I64ExtendSI32 = diff --git a/tests/Test.hs b/tests/Test.hs index f5868f7..3530abf 100644 --- a/tests/Test.hs +++ b/tests/Test.hs @@ -34,7 +34,7 @@ compile file = do main :: IO () main = do files <- Directory.listDirectory "tests/samples" - -- let files = ["data.wast"] + -- let files = ["traps.wast"] scriptTestCases <- (`mapM` files) $ \file -> do content <- LBS.readFile $ "tests/samples/" ++ file let Right script = Lexer.scanner content >>= Parser.parseScript