add some validation checks
This commit is contained in:
@@ -1,4 +1,5 @@
|
|||||||
{-# LANGUAGE NamedFieldPuns #-}
|
{-# LANGUAGE NamedFieldPuns #-}
|
||||||
|
{-# LANGUAGE DuplicateRecordFields #-}
|
||||||
|
|
||||||
module Language.Wasm.Binary (
|
module Language.Wasm.Binary (
|
||||||
dumpModule,
|
dumpModule,
|
||||||
@@ -718,7 +719,7 @@ instance Serialize LocalTypeRange where
|
|||||||
get = LocalTypeRange <$> getULEB128 <*> get
|
get = LocalTypeRange <$> getULEB128 <*> get
|
||||||
|
|
||||||
instance Serialize Function where
|
instance Serialize Function where
|
||||||
put (Function _ locals body) = do
|
put Function {locals, body} = do
|
||||||
let bs = runPut $ do
|
let bs = runPut $ do
|
||||||
putVec $ map (LocalTypeRange 1) locals
|
putVec $ map (LocalTypeRange 1) locals
|
||||||
putExpression body
|
putExpression body
|
||||||
|
|||||||
@@ -1,8 +1,71 @@
|
|||||||
|
{-# LANGUAGE NamedFieldPuns #-}
|
||||||
|
{-# LANGUAGE DuplicateRecordFields #-}
|
||||||
|
|
||||||
module Language.Wasm.Validate (
|
module Language.Wasm.Validate (
|
||||||
validate
|
validate,
|
||||||
|
isValid
|
||||||
) where
|
) where
|
||||||
|
|
||||||
import Language.Wasm.Structure
|
import Language.Wasm.Structure
|
||||||
|
import qualified Data.Set as Set
|
||||||
|
import Data.List (foldl')
|
||||||
|
import qualified Data.Text.Lazy as TL
|
||||||
|
|
||||||
validate :: Module -> Either String Module
|
data ValidationResult =
|
||||||
validate mod = Right mod
|
DuplicatedExportNames [String]
|
||||||
|
| MoreThanOneMemory
|
||||||
|
| MoreThanOneTable
|
||||||
|
| Valid
|
||||||
|
deriving (Show, Eq)
|
||||||
|
|
||||||
|
isValid :: ValidationResult -> Bool
|
||||||
|
isValid Valid = True
|
||||||
|
isValid _ = False
|
||||||
|
|
||||||
|
type Validator = Module -> ValidationResult
|
||||||
|
|
||||||
|
shouldBeAtMostOneTable :: Validator
|
||||||
|
shouldBeAtMostOneTable Module { imports, tables } =
|
||||||
|
let memImports = filter isTableImport imports in
|
||||||
|
if length memImports + length tables <= 1
|
||||||
|
then Valid
|
||||||
|
else MoreThanOneTable
|
||||||
|
where
|
||||||
|
isTableImport Import { desc = ImportTable _ } = True
|
||||||
|
isTableImport _ = False
|
||||||
|
|
||||||
|
shouldBeAtMostOneMemory :: Validator
|
||||||
|
shouldBeAtMostOneMemory Module { imports, mems } =
|
||||||
|
let memImports = filter isMemImport imports in
|
||||||
|
if length memImports + length mems <= 1
|
||||||
|
then Valid
|
||||||
|
else MoreThanOneMemory
|
||||||
|
where
|
||||||
|
isMemImport Import { desc = ImportMemory _ } = True
|
||||||
|
isMemImport _ = False
|
||||||
|
|
||||||
|
exportNamesShouldBeDifferent :: Validator
|
||||||
|
exportNamesShouldBeDifferent Module { exports } =
|
||||||
|
case foldl' go (Set.empty, []) exports of
|
||||||
|
(_, []) -> Valid
|
||||||
|
(_, dup) -> DuplicatedExportNames dup
|
||||||
|
where
|
||||||
|
go :: (Set.Set TL.Text, [String]) -> Export -> (Set.Set TL.Text, [String])
|
||||||
|
go (set, dup) (Export name _) =
|
||||||
|
if Set.member name set
|
||||||
|
then (set, show name : dup)
|
||||||
|
else (Set.insert name set, dup)
|
||||||
|
|
||||||
|
validate :: Validator
|
||||||
|
validate mod = foldl' go Valid validators
|
||||||
|
where
|
||||||
|
go :: ValidationResult -> Validator -> ValidationResult
|
||||||
|
go Valid validator = validator mod
|
||||||
|
go res _ = res
|
||||||
|
|
||||||
|
validators :: [Validator]
|
||||||
|
validators = [
|
||||||
|
shouldBeAtMostOneTable,
|
||||||
|
shouldBeAtMostOneMemory,
|
||||||
|
exportNamesShouldBeDifferent
|
||||||
|
]
|
||||||
|
|||||||
+8
-1
@@ -13,6 +13,7 @@ import qualified Language.Wasm.Lexer as Lexer
|
|||||||
import qualified Language.Wasm.Parser as Parser
|
import qualified Language.Wasm.Parser as Parser
|
||||||
import qualified Language.Wasm.Structure as Structure
|
import qualified Language.Wasm.Structure as Structure
|
||||||
import qualified Language.Wasm.Binary as Binary
|
import qualified Language.Wasm.Binary as Binary
|
||||||
|
import qualified Language.Wasm.Validate as Validate
|
||||||
|
|
||||||
import qualified Debug.Trace as Debug
|
import qualified Debug.Trace as Debug
|
||||||
|
|
||||||
@@ -42,7 +43,13 @@ main = do
|
|||||||
let Right mod' = Binary.decodeModuleLazy $ Binary.dumpModuleLazy mod
|
let Right mod' = Binary.decodeModuleLazy $ Binary.dumpModuleLazy mod
|
||||||
return $ testCase ("Dump module to binary and parse back: " ++ file) $
|
return $ testCase ("Dump module to binary and parse back: " ++ file) $
|
||||||
assertEqual "Module matched" mod mod'
|
assertEqual "Module matched" mod mod'
|
||||||
|
validationTestCases <- (`mapM` (filter (/= "import.wast") files)) $ \file -> do
|
||||||
|
content <- LBS.readFile $ "tests/samples/" ++ file
|
||||||
|
let Right mod = Parser.parseModule <$> Lexer.scanner content
|
||||||
|
return $ testCase ("Validate module: " ++ file) $
|
||||||
|
assertBool "Module matched" $ Validate.isValid $ Validate.validate mod
|
||||||
defaultMain $ testGroup "tests" [
|
defaultMain $ testGroup "tests" [
|
||||||
testGroup "Syntax parsing" syntaxTestCases,
|
testGroup "Syntax parsing" syntaxTestCases,
|
||||||
testGroup "Binary format" binaryTestCases
|
testGroup "Binary format" binaryTestCases,
|
||||||
|
testGroup "Validation" validationTestCases
|
||||||
]
|
]
|
||||||
|
|||||||
Reference in New Issue
Block a user