typecheck nested patterns and clause expressions

This commit is contained in:
milner committed 2020-01-22 12:00:00 +00:00
1 parent fb0fd5431c
commit 67f24dd845
1 file changed
+222
+222
View File
@@ -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''')