diff --git a/src/Language/Wasm/Validate.hs b/src/Language/Wasm/Validate.hs index ee7e5ae..06ebf78 100644 --- a/src/Language/Wasm/Validate.hs +++ b/src/Language/Wasm/Validate.hs @@ -10,9 +10,9 @@ module Language.Wasm.Validate ( import Language.Wasm.Structure import qualified Data.Set as Set -import Data.List (foldl', isPrefixOf) +import Data.List (foldl') import qualified Data.Text.Lazy as TL -import Data.Maybe (fromMaybe, maybeToList) +import Data.Maybe (fromMaybe, maybeToList, catMaybes) import Data.Monoid ((<>)) import Numeric.Natural (Natural) @@ -86,18 +86,23 @@ isArrowMatch (f `Arrow` t) ( f' `Arrow` t') = isEndMatch f f' && isEndMatch t t' where isEndMatch :: End -> End -> Bool isEndMatch (Any:l) (Any:r) = - let leftTail = takeWhile (/= Any) $ reverse l in - let rightTail = takeWhile (/= Any) $ reverse r in - leftTail `isPrefixOf` rightTail || rightTail `isPrefixOf` leftTail + let (leftTail, rightTail) = unzip $ zip (takeWhile (/= Any) $ reverse l) (takeWhile (/= Any) $ reverse r) in + isEndMatch (reverse leftTail) (reverse rightTail) isEndMatch (Any:l) r = - let leftTail = takeWhile (/= Any) $ reverse l in - let rightTail = takeWhile (/= Any) $ reverse r in - leftTail `isPrefixOf` rightTail || rightTail `isPrefixOf` leftTail + let (leftTail, rightTail) = unzip $ zip (takeWhile (/= Any) $ reverse l) (takeWhile (/= Any) $ reverse r) in + isEndMatch (reverse leftTail) (reverse rightTail) isEndMatch l (Any:r) = - let leftTail = takeWhile (/= Any) $ reverse l in - let rightTail = takeWhile (/= Any) $ reverse r in - leftTail `isPrefixOf` rightTail || rightTail `isPrefixOf` leftTail - isEndMatch l r = l == r + let (leftTail, rightTail) = unzip $ zip (takeWhile (/= Any) $ reverse l) (takeWhile (/= Any) $ reverse r) in + isEndMatch (reverse leftTail) (reverse rightTail) + isEndMatch (Var v:l) (x:r) = + let subst = replace (Var v) x in + isEndMatch (subst l) (subst r) + isEndMatch (x:l) (Var v:r) = + let subst = replace (Var v) x in + isEndMatch (subst l) (subst r) + isEndMatch (Val v:l) (Val v':r) = v == v' && isEndMatch l r + isEndMatch [] [] = True + isEndMatch _ _ = False data Ctx = Ctx { types :: [FuncType], @@ -368,9 +373,9 @@ unify (f `Arrow` t) (f' `Arrow` t') = unify' (f `Arrow` reverse t) (reverse f' ` let subst = replace (Var r') (Var r) in unify' (f `Arrow` t) (subst f' `Arrow` subst t') unify' (f `Arrow` (Any:_)) (_ `Arrow` t') = - return $ f `Arrow` t' - unify' (f `Arrow` _) ((Any:_) `Arrow` t') = - return $ f `Arrow` t' + return $ f `Arrow` (Any : t') + unify' (f `Arrow` _) (f'@(Any:_) `Arrow` t') = + return $ (f ++ (reverse f')) `Arrow` t' getExpressionType :: [Instruction] -> Checker Arrow getExpressionType instrs = @@ -388,17 +393,33 @@ getExpressionType instrs = go arr' rest ctxFromModule :: [ValueType] -> [Maybe ValueType] -> Maybe ValueType -> Module -> Ctx -ctxFromModule locals labels returns Module {types, functions, tables, mems, globals} = +ctxFromModule locals labels returns Module {types, functions, tables, mems, globals, imports} = + let funImports = catMaybes $ map getFuncType imports in + let tableImports = catMaybes $ map getTableType imports in + let memsImports = catMaybes $ map getMemType imports in + let globalImports = catMaybes $ map getGlobalType imports in Ctx { types, - funcs = map ((types !!) . fromIntegral . funcType) functions, - tables = map (\(Table t) -> t) tables, - mems = map (\(Memory l) -> l) mems, - globals = map (\(Global g _) -> g) globals, + funcs = funImports ++ map ((types !!) . fromIntegral . funcType) functions, + tables = tableImports ++ map (\(Table t) -> t) tables, + mems = memsImports ++ map (\(Memory l) -> l) mems, + globals = globalImports ++ map (\(Global g _) -> g) globals, locals, labels, returns } + where + getFuncType (Import _ _ (ImportFunc typeIdx)) = Just $ types !! (fromIntegral typeIdx) + getFuncType _ = Nothing + + getTableType (Import _ _ (ImportTable tableType)) = Just tableType + getTableType _ = Nothing + + getMemType (Import _ _ (ImportMemory lim)) = Just lim + getMemType _ = Nothing + + getGlobalType (Import _ _ (ImportGlobal gl)) = Just gl + getGlobalType _ = Nothing isFunctionValid :: Function -> Validator isFunctionValid Function {funcType, locals, body} mod@Module {types} = diff --git a/tests/Test.hs b/tests/Test.hs index 6a3eb4f..0aadd39 100644 --- a/tests/Test.hs +++ b/tests/Test.hs @@ -31,7 +31,8 @@ compile file = do main :: IO () main = do files <- Directory.listDirectory "tests/samples" - -- let files = ["call.wast"] + -- let files = ["call_indirect.wast", "br_table.wast", "br.wast"] + let files = ["br.wast"] -- compile "fact.wast" syntaxTestCases <- (`mapM` files) $ \file -> do content <- LBS.readFile $ "tests/samples/" ++ file diff --git a/tests/runnable/index.html b/tests/runnable/index.html index 41cf73f..e2ab999 100644 --- a/tests/runnable/index.html +++ b/tests/runnable/index.html @@ -6,6 +6,8 @@ ).then(results => { console.log('result', results) const fact = results.instance.exports['fac-iter-named-32'] + const factRec = results.instance.exports['fac-rec-named-32'] console.log(fact, fact(5)) + console.log(factRec, factRec(5)) }); diff --git a/tests/samples/fact.wast b/tests/samples/fact.wast index 258d271..6550937 100644 --- a/tests/samples/fact.wast +++ b/tests/samples/fact.wast @@ -22,6 +22,19 @@ ) ) + ;; Recursive factorial named + (func $fac-rec-named-32 (export "fac-rec-named-32") (param $n i32) (result i32) + (if (result i32) (i32.eq (get_local $n) (i32.const 0)) + (then (i32.const 1)) + (else + (i32.mul + (get_local $n) + (call $fac-rec-named-32 (i32.sub (get_local $n) (i32.const 1))) + ) + ) + ) + ) + ;; Iterative factorial (func (export "fac-iter") (param i64) (result i64) (local i64 i64)