diff --git a/src/Tessera/Check.hs b/src/Tessera/Check.hs new file mode 100644 index 0000000..d61e542 --- /dev/null +++ b/src/Tessera/Check.hs @@ -0,0 +1,222 @@ +module Tessera.Check where + +import Data.List (foldl', nub) +import qualified Data.Map.Strict as Map + +import Tessera.Env +import Tessera.Syntax + +checkProgram :: Program -> Either Diagnostic Env +checkProgram program = do + env <- buildEnv program + mapM_ (checkDeclaration env) program + return env + +checkDeclaration :: Env -> Declaration -> Either Diagnostic () +checkDeclaration env (DataDeclaration _ params constructors) = + mapM_ (checkConstructor env params) constructors +checkDeclaration env (MatchDeclaration name arguments result clauses) = do + mapM_ (checkTypeReference env (Span 0 0)) (result : map snd arguments) + mapM_ (checkClause env name (map snd arguments) result) clauses + +checkConstructor :: Env -> [String] -> Constructor -> Either Diagnostic () +checkConstructor env params (Constructor _ arguments) = + mapM_ (resolveType env params (Span 0 0)) arguments + +resolveType :: Env -> [String] -> Span -> Type -> Either Diagnostic () +resolveType env params span' ty = case ty of + TVar name + | name `elem` params -> Right () + | otherwise -> Left (Diagnostic span' ("unbound type variable " ++ name)) + TName name arguments -> do + definition <- case lookupType env name of + Just d -> Right d + Nothing -> Left (Diagnostic span' ("unknown type " ++ name)) + if length arguments /= length (typeDefinitionParams definition) + then Left (Diagnostic span' ("type " ++ name ++ " expects " ++ show (length (typeDefinitionParams definition)) ++ " arguments")) + else mapM_ (resolveType env params span') arguments + TTuple elements -> mapM_ (resolveType env params span') elements + +checkTypeReference :: Env -> Span -> Type -> Either Diagnostic () +checkTypeReference env span' ty = case ty of + TVar _ -> Right () + TName name arguments -> do + definition <- case lookupType env name of + Just d -> Right d + Nothing -> Left (Diagnostic span' ("unknown type " ++ name)) + if length arguments /= length (typeDefinitionParams definition) + then Left (Diagnostic span' ("type " ++ name ++ " expects " ++ show (length (typeDefinitionParams definition)) ++ " arguments")) + else mapM_ (checkTypeReference env span') arguments + TTuple elements -> mapM_ (checkTypeReference env span') elements + +checkClause :: Env -> String -> [Type] -> Type -> Clause -> Either Diagnostic () +checkClause env matchName argumentTypes result (Clause span' patterns body) = do + if length patterns /= length argumentTypes + then Left (Diagnostic span' ("clause for " ++ matchName ++ " has " ++ show (length patterns) ++ " patterns but the match declares " ++ show (length argumentTypes) ++ " arguments")) + else do + bindings <- foldl' mergeBindings (Right Map.empty) (zipWith (checkPattern env) argumentTypes patterns) + bodyType <- fst <$> inferExpr env bindings 0 body + subst <- unify emptySubst (applySubst emptySubst result) (applySubst emptySubst bodyType) + let _ = subst + return () + +mergeBindings :: Either Diagnostic (Map.Map String Type) -> Either Diagnostic (Map.Map String Type) -> Either Diagnostic (Map.Map String Type) +mergeBindings (Left e) _ = Left e +mergeBindings _ (Left e) = Left e +mergeBindings (Right a) (Right b) = + case [k | k <- Map.keys b, Map.member k a] of + (k : _) -> Left (Diagnostic (Span 0 0) ("duplicate pattern variable " ++ k)) + [] -> Right (Map.union a b) + +checkPattern :: Env -> Type -> Pattern -> Either Diagnostic (Map.Map String Type) +checkPattern env expected pattern = case pattern of + PWild -> Right Map.empty + PVar name -> Right (Map.singleton name expected) + PInt _ -> if compatible expected (TName "Int" []) + then Right Map.empty + else Left (Diagnostic (Span 0 0) ("integer pattern against " ++ renderType expected)) + PBool _ -> if compatible expected (TName "Bool" []) + then Right Map.empty + else Left (Diagnostic (Span 0 0) ("boolean pattern against " ++ renderType expected)) + PTuple patterns -> case expected of + TTuple types | length types == length patterns -> + foldl' mergeBindings (Right Map.empty) (zipWith (checkPattern env) types patterns) + _ -> Left (Diagnostic (Span 0 0) ("tuple pattern against " ++ renderType expected)) + PCon name patterns -> case lookupConstructor env name of + Nothing -> Left (Diagnostic (Span 0 0) ("unknown constructor " ++ name)) + Just (typeName, argumentTypes) -> do + typeArguments <- case expected of + TName other arguments | other == typeName -> Right arguments + TVar _ -> Right (map (TVar . freshName) [0 .. length argumentTypes - 1]) + _ -> Left (Diagnostic (Span 0 0) ("constructor " ++ name ++ " pattern against " ++ renderType expected)) + definition <- case lookupType env typeName of + Just d -> Right d + Nothing -> Left (Diagnostic (Span 0 0) ("unknown type " ++ typeName)) + let instantiated = instantiate (typeDefinitionParams definition) typeArguments argumentTypes + if length instantiated /= length patterns + then Left (Diagnostic (Span 0 0) ("constructor " ++ name ++ " expects " ++ show (length instantiated) ++ " arguments")) + else foldl' mergeBindings (Right Map.empty) (zipWith (checkPattern env) instantiated patterns) + +freshName :: Int -> String +freshName n = "?" ++ show n + +renderType :: Type -> String +renderType ty = case ty of + TVar name -> name + TName name [] -> name + TName name arguments -> name ++ " " ++ unwords (map renderType arguments) + TTuple elements -> "(" ++ commaJoin (map renderType elements) ++ ")" + +commaJoin :: [String] -> String +commaJoin = foldr1 (\a b -> a ++ ", " ++ b) + +compatible :: Type -> Type -> Bool +compatible (TVar _) _ = True +compatible _ (TVar _) = True +compatible a b = typeEquals a b + +data Subst = Subst (Map.Map String Type) + +emptySubst :: Subst +emptySubst = Subst Map.empty + +applySubst :: Subst -> Type -> Type +applySubst (Subst mapping) ty = case ty of + TVar name -> Map.findWithDefault (TVar name) name mapping + TName name arguments -> TName name (map (applySubst (Subst mapping)) arguments) + TTuple elements -> TTuple (map (applySubst (Subst mapping)) elements) + +composeSubst :: Subst -> Subst -> Subst +composeSubst (Subst first) (Subst second) = + Subst (Map.union (Map.map (applySubst (Subst first)) second) first) + +unify :: Subst -> Type -> Type -> Either Diagnostic Subst +unify substitution left right = + let left' = applySubst substitution left + right' = applySubst substitution right + in case (left', right') of + (TVar a, TVar b) | a == b -> Right substitution + (TVar a, other) -> Right (composeSubst (Subst (Map.singleton a other)) substitution) + (other, TVar b) -> Right (composeSubst (Subst (Map.singleton b other)) substitution) + (TName a as, TName b bs) + | a == b && length as == length bs -> + foldl' (\acc (x, y) -> acc >>= \s -> unify s x y) (Right substitution) (zip as bs) + (TTuple as, TTuple bs) + | length as == length bs -> + foldl' (\acc (x, y) -> acc >>= \s -> unify s x y) (Right substitution) (zip as bs) + _ -> Left (Diagnostic (Span 0 0) ("type mismatch between " ++ renderType left' ++ " and " ++ renderType right')) + +inferExpr :: Env -> Map.Map String Type -> Int -> Expr -> Either Diagnostic (Type, Int) +inferExpr env bindings counter expression = case expression of + EVar name -> case Map.lookup name bindings of + Just ty -> Right (ty, counter) + Nothing -> Left (Diagnostic (Span 0 0) ("unbound variable " ++ name)) + EInt _ -> Right (TName "Int" [], counter) + EBool _ -> Right (TName "Bool" [], counter) + ETuple elements -> do + results <- foldl' step (Right ([], counter)) elements + let (types, final) = results + return (TTuple types, final) + where + step (Left e) _ = Left e + step (Right (types, n)) element = do + (ty, n') <- inferExpr env bindings n element + return (types ++ [ty], n') + ECon name elements -> case lookupConstructor env name of + Nothing -> Left (Diagnostic (Span 0 0) ("unknown constructor " ++ name)) + Just (typeName, argumentTypes) -> do + definition <- case lookupType env typeName of + Just d -> Right d + Nothing -> Left (Diagnostic (Span 0 0) ("unknown type " ++ typeName)) + let params = typeDefinitionParams definition + freshVariables = [TVar (freshName (counter + i)) | i <- [0 .. length params - 1]] + instantiated = instantiate params freshVariables argumentTypes + if length elements /= length instantiated + then Left (Diagnostic (Span 0 0) ("constructor " ++ name ++ " applied to " ++ show (length elements) ++ " arguments")) + else do + substitution <- foldl' (checkArgument name) (Right emptySubst) (zip elements instantiated) + let result = applySubst substitution (TName typeName freshVariables) + return (result, counter + length params) + where + checkArgument _ (Left e) _ = Left e + checkArgument name' (Right substitution) (element, expected) = do + (actual, _) <- inferExpr env bindings counter element + unify substitution expected actual + EBin operator left right -> do + (leftType, counter') <- inferExpr env bindings counter left + (rightType, counter'') <- inferExpr env bindings counter' right + case operator of + OpAdd -> arithmetic leftType rightType counter'' + OpSub -> arithmetic leftType rightType counter'' + OpMul -> arithmetic leftType rightType counter'' + OpDiv -> arithmetic leftType rightType counter'' + OpMod -> arithmetic leftType rightType counter'' + OpLt -> arithmetic leftType rightType counter'' + OpLe -> arithmetic leftType rightType counter'' + OpGt -> arithmetic leftType rightType counter'' + OpGe -> arithmetic leftType rightType counter'' + OpEq -> do + _ <- unify emptySubst leftType rightType + return (TName "Bool" [], counter'') + OpNe -> do + _ <- unify emptySubst leftType rightType + return (TName "Bool" [], counter'') + OpAnd -> logical leftType rightType counter'' + OpOr -> logical leftType rightType counter'' + where + arithmetic leftType rightType n = do + _ <- unify emptySubst leftType (TName "Int" []) + _ <- unify emptySubst rightType (TName "Int" []) + let result = if operator `elem` [OpAdd, OpSub, OpMul, OpDiv, OpMod] then TName "Int" [] else TName "Bool" [] + return (result, n) + logical leftType rightType n = do + _ <- unify emptySubst leftType (TName "Bool" []) + _ <- unify emptySubst rightType (TName "Bool" []) + return (TName "Bool" [], n) + EIf condition yes no -> do + (conditionType, counter') <- inferExpr env bindings counter condition + _ <- unify emptySubst conditionType (TName "Bool" []) + (yesType, counter'') <- inferExpr env bindings counter' yes + (noType, counter''') <- inferExpr env bindings counter'' no + substitution <- unify emptySubst yesType noType + return (applySubst substitution yesType, counter''')