diff --git a/src/Language/Wasm/Interpreter.hs b/src/Language/Wasm/Interpreter.hs index 2f8bbba..d65978b 100644 --- a/src/Language/Wasm/Interpreter.hs +++ b/src/Language/Wasm/Interpreter.hs @@ -78,10 +78,12 @@ data Value = | VI64 Word64 | VF32 Float | VF64 Double + | VV128 ByteArray.ByteArray | RF (Maybe Natural) | RE (Maybe Natural) deriving (Eq, Show) + asInt32 :: Word32 -> Int32 asInt32 w = if w < 0x80000000 @@ -467,6 +469,7 @@ evalConstExpr _ _ [I32Const v] = return $ VI32 v evalConstExpr _ _ [I64Const v] = return $ VI64 v evalConstExpr _ _ [F32Const v] = return $ VF32 v evalConstExpr _ _ [F64Const v] = return $ VF64 v +evalConstExpr _ _ [V128Const v] = return $ VV128 v evalConstExpr _ _ [RefNull FuncRef] = return $ RF Nothing evalConstExpr _ _ [RefNull ExternRef] = return $ RE Nothing evalConstExpr inst _ [RefFunc idx] = return $ RF $ Just $ fromIntegral $ funcaddrs inst ! fromIntegral idx @@ -636,6 +639,21 @@ data EvalResult = | ReturnFn [Value] deriving (Show, Eq) +lanewise :: (Primitive.Prim i) => SimdShape -> ByteArray.ByteArray -> ByteArray.ByteArray + -> (i -> i -> i) -> ByteArray.ByteArray +lanewise shape a b op = + let count = case shape of + I8x16 -> 16 + I16x8 -> 8 + I32x4 -> 4 + I64x2 -> 2 + F32x4 -> 4 + F64x2 -> 2 + in + let proto = [0..count-1] in + ByteArray.byteArrayFromListN count + $ zipWith op (ByteArray.indexByteArray a <$> proto) (ByteArray.indexByteArray b <$> proto) + eval :: Natural -> Store -> ModuleInstance -> FunctionInstance -> [Value] -> IO (Maybe [Value]) eval 0 _ _ _ _ = return Nothing eval budget store inst FunctionInstance { funcType, moduleInstance, code = Function { localTypes, body} } args = do @@ -660,6 +678,7 @@ eval budget store inst FunctionInstance { funcType, moduleInstance, code = Funct checkValType I64 (VI64 v) = Just $ VI64 v checkValType F32 (VF32 v) = Just $ VF32 v checkValType F64 (VF64 v) = Just $ VF64 v + checkValType V128 (VV128 v) = Just $ VV128 v checkValType Func (RF v) = Just $ RF v checkValType Extern (RE v) = Just $ RE v checkValType _ _ = Nothing @@ -1079,6 +1098,7 @@ eval budget store inst FunctionInstance { funcType, moduleInstance, code = Funct step ctx (I64Const v) = return $ Done ctx { stack = VI64 v : stack ctx } step ctx (F32Const v) = return $ Done ctx { stack = VF32 v : stack ctx } step ctx (F64Const v) = return $ Done ctx { stack = VF64 v : stack ctx } + step ctx (V128Const v) = return $ Done ctx { stack = VV128 v : stack ctx } step ctx@EvalCtx{ stack = (VI32 v2:VI32 v1:rest) } (IBinOp BS32 IAdd) = return $ Done ctx { stack = VI32 (v1 + v2) : rest } step ctx@EvalCtx{ stack = (VI32 v2:VI32 v1:rest) } (IBinOp BS32 ISub) = @@ -1233,6 +1253,15 @@ eval budget store inst FunctionInstance { funcType, moduleInstance, code = Funct let half = v .&. 0xFFFFFFFF in let r = if half >= 0x80000000 then asWord64 (fromIntegral half - 0x100000000) else half in return $ Done ctx { stack = VI64 r : rest } + step ctx@EvalCtx{ stack = (VV128 v2:VV128 v1:rest) } (IBinOp (BS128 shape) IAdd) = + let r = case shape of + I8x16 -> lanewise @Word8 shape v1 v2 (+) + I16x8 -> lanewise @Word16 shape v1 v2 (+) + I32x4 -> lanewise @Word32 shape v1 v2 (+) + I64x2 -> lanewise @Word64 shape v1 v2 (+) + _ -> error "impossible due to validation" + in + return $ Done ctx { stack = VV128 r : rest } step ctx@EvalCtx{ stack = (VF32 v:rest) } (FUnOp BS32 FAbs) = return $ Done ctx { stack = VF32 (abs v) : rest } step ctx@EvalCtx{ stack = (VF32 v:rest) } (FUnOp BS32 FNeg) = diff --git a/src/Language/Wasm/Parser.y b/src/Language/Wasm/Parser.y index d4ac613..822300c 100644 --- a/src/Language/Wasm/Parser.y +++ b/src/Language/Wasm/Parser.y @@ -54,7 +54,8 @@ import Language.Wasm.Structure ( ElemType(..), Limit(..), GlobalType(..), - ValueType(..) + ValueType(..), + SimdShape(..) ) import qualified Language.Wasm.Structure as S @@ -64,7 +65,7 @@ import qualified Data.Text.Lazy as TL import qualified Data.Text.Lazy.Encoding as TLEncoding import qualified Data.Text.Lazy.Read as TLRead -import qualified Data.ByteString as BS +import qualified Data.Primitive.ByteArray as ByteArray import qualified Data.ByteString.Lazy as LBS import qualified Data.ByteString.Lazy.Char8 as LBSChar8 import Data.Maybe (fromMaybe, fromJust, isNothing, catMaybes) @@ -73,7 +74,7 @@ import Control.Monad (guard, foldM) import Control.Monad.Except (throwError) import Numeric.Natural (Natural) -import Data.Word (Word32, Word64, Word8) +import Data.Word (Word8, Word16, Word32, Word64) import Data.Bits ((.|.)) import Numeric.IEEE (infinity, nan, maxFinite) import Language.Wasm.FloatUtils (doubleToFloat, floatToWord, doubleToWord) @@ -348,6 +349,9 @@ import Language.Wasm.Lexer ( 'offset' { Lexeme _ (TKeyword "offset") } 'start' { Lexeme _ (TKeyword "start") } 'module' { Lexeme _ (TKeyword "module") } +-- simd +'i32x4.add' { Lexeme _ (TKeyword "i32x4.add") } +'i64x2.add' { Lexeme _ (TKeyword "i64x2.add") } -- script extension 'binary' { Lexeme _ (TKeyword "binary") } 'quote' { Lexeme _ (TKeyword "quote") } @@ -404,14 +408,14 @@ index :: { Index } i8 :: { Integer } : int {% - if $1 >= -(2^7) && $1 <= 2^8 + if $1 >= -(2^7) && $1 < 2^8 then Right $ fromIntegral $ if $1 >= 0 then $1 else 2^8 + $1 else Left ("I8 literal value is out of signed i8 boundaries: " ++ show $1) } i16 :: { Integer } : int {% - if $1 >= -(2^15) && $1 <= 2^16 + if $1 >= -(2^15) && $1 < 2^16 then Right $ fromIntegral $ if $1 >= 0 then $1 else 2^16 + $1 else Left ("I16 literal value is out of signed i16 boundaries: " ++ show $1) } @@ -689,6 +693,9 @@ plaininstr :: { PlainInstr } | 'i64.reinterpret_f64' { IReinterpretF BS64 } | 'f32.reinterpret_i32' { FReinterpretI BS32 } | 'f64.reinterpret_i64' { FReinterpretI BS64 } + -- simd + | 'i32x4.add' { IBinOp (BS128 I32x4) IAdd } + | 'i64x2.add' { IBinOp (BS128 I64x2) IAdd } typeuse(next) : '(' typeuse1(folded_instr_list(next), instruction_list(next)) { @@ -1285,8 +1292,6 @@ type MemoryIndex = Index type ElemIndex = Index type DataIndex = Index -data SimdShape = I8x16 | I16x8 | I32x4 | I64x2 | F32x4 | F64x2 deriving (Show, Eq) - data V128Rep = I8x16Const [Integer] | I16x8Const [Integer] @@ -1592,37 +1597,19 @@ data FunCtx = FunCtx { ctxParams :: [ParamType] } deriving (Eq, Show) -unpackWord32 :: Word32 -> [Word8] -unpackWord32 w = [ - fromIntegral $ w `rem` 0x100, - fromIntegral $ w `rem` 0x10000 `div` 0x100, - fromIntegral $ w `rem` 0x1000000 `div` 0x10000, - fromIntegral $ w `rem` 0x100000000 `div` 0x1000000] - -unpackWord64 :: Word64 -> [Word8] -unpackWord64 w = [ - fromIntegral $ w `rem` 0x100, - fromIntegral $ w `rem` 0x10000 `div` 0x100, - fromIntegral $ w `rem` 0x1000000 `div` 0x10000, - fromIntegral $ w `rem` 0x100000000 `div` 0x1000000, - fromIntegral $ w `rem` 0x10000000000 `div` 0x100000000, - fromIntegral $ w `rem` 0x1000000000000 `div` 0x10000000000, - fromIntegral $ w `rem` 0x100000000000000 `div` 0x1000000000000, - fromIntegral $ w `rem` 0x10000000000000000 `div` 0x100000000000000] - -v128RepToBytes :: V128Rep -> Either String BS.ByteString -v128RepToBytes (I8x16Const bytes) = return $ BS.pack $ fromIntegral <$> bytes +v128RepToBytes :: V128Rep -> Either String ByteArray.ByteArray +v128RepToBytes (I8x16Const bytes) = + return $ ByteArray.byteArrayFromListN 16 $ (fromIntegral :: Integer -> Word8) <$> bytes v128RepToBytes (I16x8Const words) = - let asWord8 w = [fromIntegral $ w `rem` 0x100, fromIntegral $ w `rem` 0x10000 `div` 0x100] in - return $ BS.pack $ concat $ asWord8 <$> words + return $ ByteArray.byteArrayFromListN 8 $ (fromIntegral :: Integer -> Word16) <$> words v128RepToBytes (I32x4Const dwords) = - return $ BS.pack $ concat $ unpackWord32 . integerToWord32 <$> dwords + return $ ByteArray.byteArrayFromListN 4 $ integerToWord32 <$> dwords v128RepToBytes (I64x2Const qwords) = - return $ BS.pack $ concat $ unpackWord64 . integerToWord64 <$> qwords + return $ ByteArray.byteArrayFromListN 2 $ integerToWord64 <$> qwords v128RepToBytes (F32x4Const floats) = - BS.pack . concat <$> mapM (fmap (unpackWord32 . floatToWord) . asFloat) floats + ByteArray.byteArrayFromListN 4 <$> mapM (fmap floatToWord . asFloat) floats v128RepToBytes (F64x2Const doubles) = - BS.pack . concat <$> mapM (fmap (unpackWord64 . doubleToWord) . asDouble) doubles + ByteArray.byteArrayFromListN 2 <$> mapM (fmap doubleToWord . asDouble) doubles constInstructionToValue :: Instruction -> Either String (S.Instruction Natural) constInstructionToValue (PlainInstr (I32Const v)) = return $ S.I32Const $ integerToWord32 v diff --git a/src/Language/Wasm/Script.hs b/src/Language/Wasm/Script.hs index 1c89984..728052c 100644 --- a/src/Language/Wasm/Script.hs +++ b/src/Language/Wasm/Script.hs @@ -11,6 +11,7 @@ import qualified Data.Text.Lazy.Encoding as TLEncoding import qualified Control.Monad.State as State import Control.Monad.IO.Class (liftIO) import Numeric.IEEE (identicalIEEE) +import qualified Data.Primitive.ByteArray as ByteArray import qualified Control.DeepSeq as DeepSeq import Data.Maybe (fromJust, isNothing) import Debug.Trace (trace) @@ -124,6 +125,7 @@ runScript onAssertFail script = do asArg [Struct.F32Const v] = Interpreter.VF32 v asArg [Struct.I64Const v] = Interpreter.VI64 v asArg [Struct.F64Const v] = Interpreter.VF64 v + asArg [Struct.V128Const v] = Interpreter.VV128 v asArg [Struct.RefNull Struct.FuncRef] = Interpreter.RF Nothing asArg [Struct.RefNull Struct.ExternRef] = Interpreter.RE Nothing asArg [Struct.RefExtern v] = Interpreter.RE (Just v) @@ -144,6 +146,7 @@ runScript onAssertFail script = do isValueEqual (Interpreter.VI64 v1) (Interpreter.VI64 v2) = v1 == v2 isValueEqual (Interpreter.VF32 v1) (Interpreter.VF32 v2) = (isNaN v1 && isNaN v2) || identicalIEEE v1 v2 isValueEqual (Interpreter.VF64 v1) (Interpreter.VF64 v2) = (isNaN v1 && isNaN v2) || identicalIEEE v1 v2 + isValueEqual (Interpreter.VV128 a) (Interpreter.VV128 b) = ByteArray.compareByteArrays a 0 b 0 16 == EQ isValueEqual (Interpreter.RF f1) (Interpreter.RF f2) = f1 == f2 isValueEqual (Interpreter.RE e1) (Interpreter.RE e2) = e1 == e2 isValueEqual _ _ = False diff --git a/src/Language/Wasm/Structure.hs b/src/Language/Wasm/Structure.hs index 2bef576..0b1b3a0 100644 --- a/src/Language/Wasm/Structure.hs +++ b/src/Language/Wasm/Structure.hs @@ -33,6 +33,7 @@ module Language.Wasm.Structure ( FuncType(..), ValueType(..), BlockType(..), + SimdShape(..), ParamsType, ResultType, LocalsType, @@ -53,13 +54,15 @@ module Language.Wasm.Structure ( import Numeric.Natural (Natural) import Data.Word (Word32, Word64) -import qualified Data.ByteString as BS +import qualified Data.Primitive.ByteArray as ByteArray import qualified Data.ByteString.Lazy as LBS import qualified Data.Text.Lazy as TL import Control.DeepSeq (NFData) import GHC.Generics (Generic) -data BitSize = BS32 | BS64 | BS128 deriving (Show, Eq, Generic, NFData) +data SimdShape = I8x16 | I16x8 | I32x4 | I64x2 | F32x4 | F64x2 | I128x1 deriving (Show, Eq, Generic, NFData) + +data BitSize = BS32 | BS64 | BS128 SimdShape deriving (Show, Eq, Generic, NFData) data IUnOp = IClz @@ -200,7 +203,7 @@ data Instruction index = | I64Const Word64 | F32Const Float | F64Const Double - | V128Const BS.ByteString + | V128Const ByteArray.ByteArray | IUnOp BitSize IUnOp | IBinOp BitSize IBinOp | I32Eqz diff --git a/src/Language/Wasm/Validate.hs b/src/Language/Wasm/Validate.hs index 6cfefc8..9d688cd 100644 --- a/src/Language/Wasm/Validate.hs +++ b/src/Language/Wasm/Validate.hs @@ -467,10 +467,12 @@ getInstrType _ (I32Const _) = return $ empty ==> I32 getInstrType _ (I64Const _) = return $ empty ==> I64 getInstrType _ (F32Const _) = return $ empty ==> F32 getInstrType _ (F64Const _) = return $ empty ==> F64 +getInstrType _ (V128Const _) = return $ empty ==> V128 getInstrType _ (IUnOp BS32 _) = return $ I32 ==> I32 getInstrType _ (IUnOp BS64 _) = return $ I64 ==> I64 getInstrType _ (IBinOp BS32 _) = return $ [I32, I32] ==> I32 getInstrType _ (IBinOp BS64 _) = return $ [I64, I64] ==> I64 +getInstrType _ (IBinOp (BS128 _) _) = return $ [V128, V128] ==> V128 getInstrType _ I32Eqz = return $ I32 ==> I32 getInstrType _ I64Eqz = return $ I64 ==> I32 getInstrType _ (IRelOp BS32 _) = return $ [I32, I32] ==> I32 @@ -575,6 +577,7 @@ isConstExpression ((I32Const _):rest) = isConstExpression rest isConstExpression ((I64Const _):rest) = isConstExpression rest isConstExpression ((F32Const _):rest) = isConstExpression rest isConstExpression ((F64Const _):rest) = isConstExpression rest +isConstExpression ((V128Const _):rest) = isConstExpression rest isConstExpression ((RefNull _):rest) = isConstExpression rest isConstExpression ((RefFunc _):rest) = isConstExpression rest isConstExpression ((GetGlobal idx):rest) = do