diff --git a/src/Language/Wasm/Script.hs b/src/Language/Wasm/Script.hs index 8a585b6..a737605 100644 --- a/src/Language/Wasm/Script.hs +++ b/src/Language/Wasm/Script.hs @@ -152,26 +152,29 @@ runScript onAssertFail script = do checkModuleInvalid :: Struct.Module -> IO () checkModuleInvalid _ = return () - getFailureString :: Validate.ValidationResult -> TL.Text - getFailureString (Validate.TypeMismatch _ _) = "type mismatch" - getFailureString Validate.ResultTypeDoesntMatch = "type mismatch" - getFailureString Validate.MoreThanOneMemory = "multiple memories" - getFailureString Validate.MoreThanOneTable = "multiple tables" - getFailureString Validate.LocalIndexOutOfRange = "unknown local" - getFailureString Validate.MemoryIndexOutOfRange = "unknown memory" - getFailureString Validate.TableIndexOutOfRange = "unknown table" - getFailureString Validate.FunctionIndexOutOfRange = "unknown function" - getFailureString Validate.GlobalIndexOutOfRange = "unknown global" - getFailureString Validate.LabelIndexOutOfRange = "unknown label" - getFailureString Validate.TypeIndexOutOfRange = "unknown type" - getFailureString Validate.MinMoreThanMaxInMemoryLimit = "size minimum must not be greater than maximum" - getFailureString Validate.MemoryLimitExceeded = "memory size must be at most 65536 pages (4GiB)" - getFailureString Validate.AlignmentOverflow = "alignment must not be larger than natural" - getFailureString (Validate.DuplicatedExportNames _) = "duplicate export name" - getFailureString Validate.InvalidConstantExpr = "constant expression required" - getFailureString Validate.InvalidResultArity = "invalid result arity" - -- getFailureString Validate.ImportedGlobalIsNotConst = "global is immutable" - getFailureString _ = "not implemented" + getFailureString :: Validate.ValidationResult -> [TL.Text] + getFailureString (Validate.TypeMismatch _ _) = ["type mismatch"] + getFailureString Validate.ResultTypeDoesntMatch = ["type mismatch"] + getFailureString Validate.MoreThanOneMemory = ["multiple memories"] + getFailureString Validate.MoreThanOneTable = ["multiple tables"] + getFailureString Validate.LocalIndexOutOfRange = ["unknown local"] + getFailureString Validate.MemoryIndexOutOfRange = ["unknown memory", "unknown memory 0"] + getFailureString Validate.TableIndexOutOfRange = ["unknown table", "unknown table 0"] + getFailureString Validate.FunctionIndexOutOfRange = ["unknown function", "unknown function 0"] + getFailureString Validate.GlobalIndexOutOfRange = ["unknown global"] + getFailureString Validate.LabelIndexOutOfRange = ["unknown label"] + getFailureString Validate.TypeIndexOutOfRange = ["unknown type"] + getFailureString Validate.MinMoreThanMaxInMemoryLimit = ["size minimum must not be greater than maximum"] + getFailureString Validate.MemoryLimitExceeded = ["memory size must be at most 65536 pages (4GiB)"] + getFailureString Validate.AlignmentOverflow = ["alignment", "alignment must not be larger than natural"] + getFailureString (Validate.DuplicatedExportNames _) = ["duplicate export name"] + getFailureString Validate.InvalidConstantExpr = ["constant expression required"] + getFailureString Validate.InvalidResultArity = ["invalid result arity"] + getFailureString Validate.GlobalIsImmutable = ["global is immutable"] + getFailureString Validate.ImportedGlobalIsNotConst = ["mutable globals cannot be imported"] + getFailureString Validate.ExportedGlobalIsNotConst = ["mutable globals cannot be exported"] + getFailureString Validate.InvalidStartFunctionType = ["start function"] + getFailureString r = [TL.concat ["not implemented ", (TL.pack $ show r)]] runAssert :: ScriptState -> Assertion -> IO () runAssert st assert@(AssertReturn action expected) = do @@ -186,7 +189,7 @@ runScript onAssertFail script = do case Validate.validate m of Validate.Valid -> onAssertFail "Invalid module pass validation" assert reason -> - if getFailureString reason == failureString + if failureString `elem` getFailureString reason then return () else let msg = "Module invalid for other reason. Expected " diff --git a/src/Language/Wasm/Validate.hs b/src/Language/Wasm/Validate.hs index 91a86a6..3edeb0f 100644 --- a/src/Language/Wasm/Validate.hs +++ b/src/Language/Wasm/Validate.hs @@ -12,7 +12,7 @@ import Language.Wasm.Structure import qualified Data.Set as Set import Data.List (foldl') import qualified Data.Text.Lazy as TL -import Data.Maybe (fromMaybe, maybeToList, catMaybes, isNothing) +import Data.Maybe (fromMaybe, maybeToList, catMaybes) import Data.Monoid ((<>)) import Numeric.Natural (Natural) @@ -43,6 +43,8 @@ data ValidationResult = | InvalidConstantExpr | InvalidStartFunctionType | ImportedGlobalIsNotConst + | ExportedGlobalIsNotConst + | GlobalIsImmutable | Valid deriving (Show, Eq) @@ -152,6 +154,10 @@ asType :: GlobalType -> VType asType (Const v) = Val v asType (Mut v) = Val v +shouldBeMut :: GlobalType -> Checker () +shouldBeMut (Mut _) = return () +shouldBeMut (Const v) = throwError GlobalIsImmutable + getLabel :: LabelIndex -> Checker (Maybe ValueType) getLabel lbl = do Ctx { labels } <- ask @@ -202,9 +208,7 @@ getInstrType (BrIf lbl) = do getInstrType (BrTable lbls lbl) = do r <- getLabel lbl rs <- mapM getLabel lbls - -- this check for equality doesn't match the spec, - -- but a reference compiler does the same - if all (\r' -> (isNothing r && isNothing r') || r' == r || isNothing r') rs + if all (== r) rs then return $ ([Any] ++ (map Val $ maybeToList r) ++ [Val I32]) ==> Any else throwError ResultTypeDoesntMatch getInstrType Return = do @@ -218,7 +222,7 @@ getInstrType (CallIndirect sign) = do if length tables < 1 then throwError TableIndexOutOfRange else do - Arrow from to <- maybeToEither FunctionIndexOutOfRange $ asArrow <$> types !? sign + Arrow from to <- maybeToEither TypeIndexOutOfRange $ asArrow <$> types !? sign return $ (from ++ [Val I32]) ==> to getInstrType Drop = do var <- freshVar @@ -240,11 +244,12 @@ getInstrType (TeeLocal local) = do return $ Val t ==> Val t getInstrType (GetGlobal global) = do Ctx { globals } <- ask - t <- maybeToEither LocalIndexOutOfRange $ asType <$> globals !? global + t <- maybeToEither GlobalIndexOutOfRange $ asType <$> globals !? global return $ empty ==> t getInstrType (SetGlobal global) = do Ctx { globals } <- ask - t <- maybeToEither LocalIndexOutOfRange $ asType <$> globals !? global + t <- maybeToEither GlobalIndexOutOfRange $ asType <$> globals !? global + shouldBeMut $ globals !! fromIntegral global return $ t ==> empty getInstrType (I32Load memarg) = do checkMemoryInstr 4 memarg @@ -401,7 +406,7 @@ isConstExpression ((F64Const _):rest) = isConstExpression rest isConstExpression ((GetGlobal idx):rest) = do Ctx {globals, importedGlobals} <- ask if importedGlobals <= idx - then throwError InvalidConstantExpr + then throwError GlobalIndexOutOfRange else return () case globals !! fromIntegral idx of Const _ -> isConstExpression rest @@ -570,7 +575,7 @@ startShouldBeValid m@Module { start = Just (StartFunction idx) } = let i = fromIntegral idx in if length types > i then if FuncType [] [] == types !! i then Valid else InvalidStartFunctionType - else TableIndexOutOfRange + else FunctionIndexOutOfRange exportsShouldBeValid :: Validator exportsShouldBeValid Module { exports, imports, functions, mems, tables, globals } = @@ -589,7 +594,17 @@ exportsShouldBeValid Module { exports, imports, functions, mems, tables, globals isExportValid (Export _ (ExportMemory memIdx)) = if fromIntegral memIdx < length memImports + length mems then Valid else MemoryIndexOutOfRange isExportValid (Export _ (ExportGlobal globalIdx)) = - if fromIntegral globalIdx < length globalImports + length globals then Valid else GlobalIndexOutOfRange + if fromIntegral globalIdx < length globalImports + length globals + then ( + if fromIntegral globalIdx >= length globalImports + then ( + case globals !! (fromIntegral globalIdx - length globalImports) of + (Global (Mut _) _) -> ExportedGlobalIsNotConst + _ -> Valid + ) + else Valid + ) + else GlobalIndexOutOfRange areExportNamesUnique :: ValidationResult areExportNamesUnique = diff --git a/tests/Test.hs b/tests/Test.hs index 9ffdda3..d212491 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 = ["br_if.wast"] + -- let files = ["data.wast"] scriptTestCases <- (`mapM` files) $ \file -> do content <- LBS.readFile $ "tests/samples/" ++ file let Right script = Parser.parseScript <$> Lexer.scanner content