diff --git a/src/Language/Wasm/Interpreter.hs b/src/Language/Wasm/Interpreter.hs index 899db49..8f0835f 100644 --- a/src/Language/Wasm/Interpreter.hs +++ b/src/Language/Wasm/Interpreter.hs @@ -982,14 +982,14 @@ eval budget store FunctionInstance { funcType, moduleInstance, code = Function { let TableInstance { items } = tableInstances store ! tableAddr let dst = fromIntegral offset let val = case ref of - RE extRef -> extRef - RF fnRef -> fnRef + RE extRef -> fromIntegral <$> extRef + RF fnRef -> (funcaddrs moduleInstance !) . fromIntegral <$> fnRef v -> error "Impossible due to validation" els <- readIORef items if dst >= MVector.length els then return Trap else do - MVector.unsafeWrite els dst (fromIntegral <$> val) + MVector.unsafeWrite els dst val return $ Done ctx { stack = rest } step ctx@EvalCtx{ stack = (VI32 offset:rest) } (TableGet tableIdx) = do let tableAddr = tableaddrs moduleInstance ! fromIntegral tableIdx diff --git a/src/Language/Wasm/Script.hs b/src/Language/Wasm/Script.hs index 7872005..44cc990 100644 --- a/src/Language/Wasm/Script.hs +++ b/src/Language/Wasm/Script.hs @@ -180,7 +180,7 @@ runScript onAssertFail script = do getFailureString (Validate.LocalIndexOutOfRange idx) = ["unknown local", "unknown local " <> TL.pack (show idx)] getFailureString (Validate.MemoryIndexOutOfRange idx) = ["unknown memory", "unknown memory " <> TL.pack (show idx)] getFailureString (Validate.TableIndexOutOfRange idx) = ["unknown table", "unknown table " <> TL.pack (show idx)] - getFailureString Validate.FunctionIndexOutOfRange = ["unknown function", "unknown function 0"] + getFailureString (Validate.FunctionIndexOutOfRange idx) = ["unknown function", "unknown function " <> TL.pack (show idx)] getFailureString (Validate.GlobalIndexOutOfRange idx) = ["unknown global", "unknown global " <> TL.pack (show idx)] getFailureString Validate.LabelIndexOutOfRange = ["unknown label"] getFailureString Validate.TypeIndexOutOfRange = ["unknown type"] @@ -194,6 +194,7 @@ runScript onAssertFail script = do getFailureString Validate.InvalidStartFunctionType = ["start function"] getFailureString Validate.InvalidTableType = ["size minimum must not be greater than maximum"] getFailureString (Validate.ElemIndexOutOfRange idx) = ["unknown elem segment " <> TL.pack (show idx)] + getFailureString (Validate.UndeclaredFunctionRef _) = ["undeclared function reference"] getFailureString r = [TL.concat ["not implemented ", TL.pack $ show r]] printFailedAssert :: String -> Assertion -> AssertM () diff --git a/src/Language/Wasm/Validate.hs b/src/Language/Wasm/Validate.hs index 315e6bd..df54e5d 100644 --- a/src/Language/Wasm/Validate.hs +++ b/src/Language/Wasm/Validate.hs @@ -32,7 +32,7 @@ data ValidationError = | MemoryLimitExceeded | AlignmentOverflow | MoreThanOneMemory - | FunctionIndexOutOfRange + | FunctionIndexOutOfRange Natural | TableIndexOutOfRange Natural | MemoryIndexOutOfRange Natural | LocalIndexOutOfRange Natural @@ -47,6 +47,7 @@ data ValidationError = | InvalidConstantExpr | InvalidStartFunctionType | GlobalIsImmutable + | UndeclaredFunctionRef Natural deriving (Show, Eq) type ValidationResult = Either ValidationError () @@ -134,7 +135,8 @@ data Ctx = Ctx { locals :: [ValueType], labels :: [[ValueType]], returns :: [ValueType], - importedGlobals :: Natural + importedGlobals :: Natural, + refs :: Set.Set Natural } deriving (Show, Eq) type Checker = ReaderT Ctx (Except ValidationError) @@ -250,7 +252,7 @@ getInstrType Return = do return $ (Any : (map Val returns)) ==> Any getInstrType (Call fun) = do Ctx { funcs } <- ask - maybeToEither FunctionIndexOutOfRange $ asArrow <$> funcs !? fun + maybeToEither (FunctionIndexOutOfRange fun) $ asArrow <$> funcs !? fun getInstrType (CallIndirect tableIdx sign) = do Ctx { types, tables } <- ask if length tables <= fromIntegral tableIdx @@ -271,10 +273,13 @@ getInstrType RefIsNull = do var <- freshVar return $ var ==> Val I32 getInstrType (RefFunc funIdx) = do - Ctx { funcs } <- ask + Ctx { funcs, refs } <- ask if fromIntegral funIdx < length funcs - then return $ empty ==> Val Func - else throwError FunctionIndexOutOfRange + then do + unless (Set.member funIdx refs) $ + throwError $ UndeclaredFunctionRef $ fromIntegral funIdx + return $ empty ==> Val Func + else throwError $ FunctionIndexOutOfRange $ fromIntegral funIdx getInstrType (GetLocal local) = do Ctx { locals } <- ask t <- maybeToEither (LocalIndexOutOfRange local) $ locals !? local @@ -523,7 +528,7 @@ getFuncTypes Module {types, functions, imports} = getFuncType _ = Nothing ctxFromModule :: [ValueType] -> [[ValueType]] -> [ValueType] -> Module -> Ctx -ctxFromModule locals labels returns m@Module {types, tables, mems, globals, imports, elems} = +ctxFromModule locals labels returns m@Module {types, tables, mems, globals, imports, elems, exports} = let tableImports = catMaybes $ map getTableType imports in let memsImports = catMaybes $ map getMemType imports in let globalImports = catMaybes $ map getGlobalType imports in @@ -537,7 +542,10 @@ ctxFromModule locals labels returns m@Module {types, tables, mems, globals, impo locals, labels, returns, - importedGlobals = fromIntegral $ length globalImports + importedGlobals = fromIntegral $ length globalImports, + refs = Set.unions $ map getElemRefs elems + ++ map getGlobalRefs globals + ++ map getExportRefs exports } where getTableType (Import _ _ (ImportTable tableType)) = Just tableType @@ -549,6 +557,19 @@ ctxFromModule locals labels returns m@Module {types, tables, mems, globals, impo getGlobalType (Import _ _ (ImportGlobal gl)) = Just gl getGlobalType _ = Nothing + getElemRefs ElemSegment{ elemType = FuncRef, elements} = + foldl extractRef Set.empty elements + where + extractRef refs [RefFunc idx] = Set.insert idx refs + extractRef refs _ = refs + getElemRefs _ = Set.empty + + getGlobalRefs Global {initializer = [RefFunc idx]} = Set.singleton idx + getGlobalRefs _ = Set.empty + + getExportRefs Export {desc = ExportFunc idx} = Set.singleton idx + getExportRefs _ = Set.empty + isFunctionValid :: Function -> Validator isFunctionValid Function {funcType, localTypes = locals, body} mod@Module {types} = if fromIntegral funcType < length types @@ -664,7 +685,7 @@ startShouldBeValid m@Module { start = Just (StartFunction idx) } = let i = fromIntegral idx in if length types > i then if FuncType [] [] == types !! i then return () else Left InvalidStartFunctionType - else Left FunctionIndexOutOfRange + else Left $ FunctionIndexOutOfRange $ fromIntegral i exportsShouldBeValid :: Validator exportsShouldBeValid Module { exports, imports, functions, mems, tables, globals } = @@ -677,7 +698,7 @@ exportsShouldBeValid Module { exports, imports, functions, mems, tables, globals isExportValid :: Export -> ValidationResult isExportValid (Export _ (ExportFunc funIdx)) = - if fromIntegral funIdx < length funcImports + length functions then return () else Left FunctionIndexOutOfRange + if fromIntegral funIdx < length funcImports + length functions then return () else Left (FunctionIndexOutOfRange funIdx) isExportValid (Export _ (ExportTable tableIdx)) = if fromIntegral tableIdx < length tableImports + length tables then return () else Left (TableIndexOutOfRange tableIdx) isExportValid (Export _ (ExportMemory memIdx)) = diff --git a/tests/Test.hs b/tests/Test.hs index 0c0b193..45f8ec1 100644 --- a/tests/Test.hs +++ b/tests/Test.hs @@ -19,7 +19,7 @@ main = do files <- filter (not . List.isPrefixOf "simd") . filter (List.isSuffixOf ".wast") <$> Directory.listDirectory "tests/spec" - -- let files = ["table_get.wast"] + -- let files = ["ref_func.wast"] scriptTestCases <- (`mapM` files) $ \file -> do test <- LBS.readFile ("tests/spec/" ++ file) return $ testCase file $ do