typecheck nested patterns and clause expressions
This commit is contained in:
1 file changed
+222
@@ -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''')
|
||||
Reference in new issue
Block a user