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