{-# LANGUAGE NoImplicitPrelude #-} {-# LANGUAGE NamedFieldPuns #-} module Felix.Checking.Datatype ( CheckedDatatype , DatatypeValidationError , datatypeValidationErrorLocation , renderDatatypeValidationError , prepareCheckedDatatype , checkedDatatypeHeadSymbol , checkedDatatypeConstructorSymbols , CheckedDatatypeClauseView(..) , CheckedDatatypePremiseView(..) , checkedDatatypeClauseViews , checkedDatatypeGeneratedFacts ) where import Base import Felix.Report.Location import Felix.Syntax.Internal import Felix.Syntax.Lexicon import Data.List qualified as List import Data.List.NonEmpty qualified as NonEmpty import Data.Set qualified as Set import Data.Text qualified as Text -- | A datatype whose premise domains have been canonicalized and validated. data CheckedDatatype = CheckedDatatype { checkedDatatypeHead :: !FunctionSymbol , checkedDatatypeClauses :: !(NonEmpty CheckedDatatypeClause) } data CheckedDatatypeClause = CheckedDatatypeClause { checkedDatatypeClauseConstructor :: !FunctionSymbol , checkedDatatypeClauseConstructorArgs :: ![VarSymbol] , checkedDatatypeClausePremises :: ![CheckedDatatypePremise] } -- Direct premises retain their canonical term for generated-fact locations. data CheckedDatatypePremise = RecursiveDatatypePremise !VarSymbol !Expr | NonRecursiveDatatypePremise !VarSymbol !Expr data DatatypeValidationError = InvalidDatatype !Text | InvalidDatatypePremise !Location !Text deriving (Show, Eq) datatypeValidationErrorLocation :: DatatypeValidationError -> Maybe Location datatypeValidationErrorLocation = \case InvalidDatatype _message -> Nothing InvalidDatatypePremise location _message -> Just location renderDatatypeValidationError :: DatatypeValidationError -> Text renderDatatypeValidationError = \case InvalidDatatype message -> message InvalidDatatypePremise _location message -> message data CanonicalDatatypeClause = CanonicalDatatypeClause { canonicalDatatypeClauseConstructor :: !SymbolPattern , canonicalDatatypeClausePremises :: ![(VarSymbol, Location, Expr)] } -- | Read-only normalized clause data for exact typed lowering. data CheckedDatatypeClauseView = CheckedDatatypeClauseView { checkedDatatypeClauseViewConstructor :: !FunctionSymbol , checkedDatatypeClauseViewArguments :: ![VarSymbol] , checkedDatatypeClauseViewPremises :: ![CheckedDatatypePremiseView] } deriving (Show, Eq) data CheckedDatatypePremiseView = CheckedRecursiveDatatypePremise !VarSymbol !Expr | CheckedNonRecursiveDatatypePremise !VarSymbol !Expr deriving (Show, Eq) -- | Canonicalize every premise domain once, then validate the declaration. prepareCheckedDatatype :: Monad m => (Expr -> m Expr) -> Datatype -> m (Either DatatypeValidationError CheckedDatatype) prepareCheckedDatatype canonicalizeDomain Datatype{datatypeHead, datatypeClauses} = do canonicalClauses <- traverse canonicalizeClause datatypeClauses pure (validateDatatype datatypeHead canonicalClauses) where canonicalizeClause DatatypeClause { datatypeClauseConstructor , datatypeClausePremises } = do premises <- traverse (\(var, domain) -> (\canonicalDomain -> (var, exprLocation domain, canonicalDomain)) <$> canonicalizeDomain domain) datatypeClausePremises pure CanonicalDatatypeClause { canonicalDatatypeClauseConstructor = datatypeClauseConstructor , canonicalDatatypeClausePremises = premises } validateDatatype :: SymbolPattern -> NonEmpty CanonicalDatatypeClause -> Either DatatypeValidationError CheckedDatatype validateDatatype (SymbolPattern datatypeSymbol datatypeArgs) datatypeClauses = do require (null datatypeArgs) "datatype head must be nullary" let constructorSymbols = [ constructorSymbol | CanonicalDatatypeClause { canonicalDatatypeClauseConstructor = SymbolPattern constructorSymbol _ } <- NonEmpty.toList datatypeClauses ] duplicateConstructors = duplicateDatatypeConstructorSymbols constructorSymbols require (Set.null duplicateConstructors) ( "datatype constructor patterns must be distinct: " <> formatDatatypeConstructors duplicateConstructors ) checkedDatatypeClauses <- traverse (validateDatatypeClause datatypeSymbol) datatypeClauses pure CheckedDatatype { checkedDatatypeHead = datatypeSymbol , checkedDatatypeClauses } validateDatatypeClause :: FunctionSymbol -> CanonicalDatatypeClause -> Either DatatypeValidationError CheckedDatatypeClause validateDatatypeClause datatypeSymbol CanonicalDatatypeClause { canonicalDatatypeClauseConstructor = SymbolPattern constructorSymbol constructorArgs , canonicalDatatypeClausePremises } = do require (constructorSymbol /= ApplySymbol) "datatype constructor cannot be function application" require (constructorSymbol /= datatypeSymbol) "datatype constructor cannot reuse the datatype symbol" let duplicateConstructorArgs = duplicateVars constructorArgs require (Set.null duplicateConstructorArgs) ( "datatype constructor arguments must be linear: " <> formatVars duplicateConstructorArgs ) let premiseVars = (\(variable, _location, _domain) -> variable) <$> canonicalDatatypeClausePremises duplicatePremiseVars = duplicateVars premiseVars require (Set.null duplicatePremiseVars) ( "datatype premise variables must be linear: " <> formatVars duplicatePremiseVars ) let missingPremises = Set.fromList constructorArgs `Set.difference` Set.fromList premiseVars require (Set.null missingPremises) ( "datatype constructor argument(s) missing premise: " <> formatVars missingPremises ) checkedDatatypeClausePremises <- traverse (validateDatatypePremise datatypeSymbol) canonicalDatatypeClausePremises pure CheckedDatatypeClause { checkedDatatypeClauseConstructor = constructorSymbol , checkedDatatypeClauseConstructorArgs = constructorArgs , checkedDatatypeClausePremises = checkedDatatypeClausePremises } validateDatatypePremise :: FunctionSymbol -> (VarSymbol, Location, Expr) -> Either DatatypeValidationError CheckedDatatypePremise validateDatatypePremise datatypeSymbol (var, location, domain) = do let domainFreeVars = freeVars domain requirePremise location (Set.null domainFreeVars) ( "datatype premise domains must be closed terms: " <> formatVars domainFreeVars ) if equivalent domain datatypeCarrier then Right (RecursiveDatatypePremise var domain) else do requirePremise location ( SymbolMixfix datatypeSymbol `Set.notMember` mentionedSymbols domain ) "datatype recursive premise must be direct" Right (NonRecursiveDatatypePremise var domain) where datatypeCarrier = TermSymbol Nowhere (SymbolMixfix datatypeSymbol) [] require :: Bool -> Text -> Either DatatypeValidationError () require condition message | condition = Right () | otherwise = Left (InvalidDatatype message) requirePremise :: Location -> Bool -> Text -> Either DatatypeValidationError () requirePremise location condition message | condition = Right () | otherwise = Left (InvalidDatatypePremise location message) checkedDatatypeHeadSymbol :: CheckedDatatype -> Symbol checkedDatatypeHeadSymbol = SymbolMixfix . checkedDatatypeHead checkedDatatypeConstructorSymbols :: CheckedDatatype -> NonEmpty Symbol checkedDatatypeConstructorSymbols = fmap (SymbolMixfix . checkedDatatypeClauseConstructor) . checkedDatatypeClauses checkedDatatypeClauseViews :: CheckedDatatype -> NonEmpty CheckedDatatypeClauseView checkedDatatypeClauseViews = fmap clauseView . checkedDatatypeClauses where clauseView clause = CheckedDatatypeClauseView { checkedDatatypeClauseViewConstructor = checkedDatatypeClauseConstructor clause , checkedDatatypeClauseViewArguments = checkedDatatypeClauseConstructorArgs clause , checkedDatatypeClauseViewPremises = premiseView <$> checkedDatatypeClausePremises clause } premiseView = \case RecursiveDatatypePremise variable domain -> CheckedRecursiveDatatypePremise variable domain NonRecursiveDatatypePremise variable domain -> CheckedNonRecursiveDatatypePremise variable domain checkedDatatypeGeneratedFacts :: CheckedDatatype -> NonEmpty (Marker, Formula) checkedDatatypeGeneratedFacts = datatypeFacts datatypeFacts :: CheckedDatatype -> NonEmpty (Marker, Formula) datatypeFacts datatype = appendList (datatypeIntroFacts datatype) ( datatypeDistinctFacts datatype <> datatypeInjectiveFacts datatype <> [ ( datatypeCasesMarker datatype , datatypeCasesFormula datatype ) , ( datatypeInductMarker datatype , datatypeInductFormula datatype ) ] ) where appendList (first :| rest) trailing = first :| (rest <> trailing) datatypeIntroFacts :: CheckedDatatype -> NonEmpty (Marker, Formula) datatypeIntroFacts datatype = fmap (\clause -> ( datatypeIntroMarker datatype clause , datatypeIntroFormula datatype clause )) (checkedDatatypeClauses datatype) datatypeDistinctFacts :: CheckedDatatype -> [(Marker, Formula)] datatypeDistinctFacts datatype = [ ( datatypeDistinctMarker datatype leftClause rightClause , datatypeDistinctFormula leftClause rightClause ) | (leftClause, rightClause) <- unorderedPairs (NonEmpty.toList (checkedDatatypeClauses datatype)) ] datatypeInjectiveFacts :: CheckedDatatype -> [(Marker, Formula)] datatypeInjectiveFacts datatype = [ ( datatypeInjectiveMarker datatype clause , datatypeInjectiveFormula clause ) | clause <- NonEmpty.toList (checkedDatatypeClauses datatype) , not (null (checkedDatatypeClauseConstructorArgs clause)) ] datatypeIntroMarker :: CheckedDatatype -> CheckedDatatypeClause -> Marker datatypeIntroMarker datatype clause = Marker ( datatypeFactBase datatype <> "_" <> datatypeClauseSymbolText clause <> "_intro" ) datatypeDistinctMarker :: CheckedDatatype -> CheckedDatatypeClause -> CheckedDatatypeClause -> Marker datatypeDistinctMarker datatype leftClause rightClause = Marker ( datatypeFactBase datatype <> "_" <> datatypeClauseSymbolText leftClause <> "_" <> datatypeClauseSymbolText rightClause <> "_distinct" ) datatypeInjectiveMarker :: CheckedDatatype -> CheckedDatatypeClause -> Marker datatypeInjectiveMarker datatype clause = Marker ( datatypeFactBase datatype <> "_" <> datatypeClauseSymbolText clause <> "_injective" ) datatypeCasesMarker :: CheckedDatatype -> Marker datatypeCasesMarker datatype = Marker (datatypeFactBase datatype <> "_cases") datatypeInductMarker :: CheckedDatatype -> Marker datatypeInductMarker datatype = Marker (datatypeFactBase datatype <> "_induct") datatypeFactBase :: CheckedDatatype -> Text datatypeFactBase = functionSymbolText . checkedDatatypeHead datatypeIntroFormula :: CheckedDatatype -> CheckedDatatypeClause -> Formula datatypeIntroFormula datatype clause = forallIfNeeded premiseVars (impliesFrom premises conclusion) where premiseVars = datatypeClausePremiseVars clause premises = datatypeClausePremiseFormulas clause conclusion = datatypeClauseResultFormula datatype clause datatypeDistinctFormula :: CheckedDatatypeClause -> CheckedDatatypeClause -> Formula datatypeDistinctFormula leftClause rightClause = forallIfNeeded (leftVars <> rightVars) (NotEquals Nowhere leftTerm rightTerm) where leftVars = checkedDatatypeClauseConstructorArgs leftClause rightVars = renameDatatypeVars (Set.fromList leftVars) (checkedDatatypeClauseConstructorArgs rightClause) leftTerm = datatypeClauseTerm leftClause (TermVar <$> leftVars) rightTerm = datatypeClauseTerm rightClause (TermVar <$> rightVars) datatypeInjectiveFormula :: CheckedDatatypeClause -> Formula datatypeInjectiveFormula clause = forallIfNeeded (leftVars <> rightVars) ( Equals Nowhere leftTerm rightTerm `Implies` makeConjunction equalities ) where leftVars = checkedDatatypeClauseConstructorArgs clause rightVars = renameDatatypeVars (Set.fromList leftVars) leftVars leftTerm = datatypeClauseTerm clause (TermVar <$> leftVars) rightTerm = datatypeClauseTerm clause (TermVar <$> rightVars) equalities = zipWith (\leftVar rightVar -> Equals Nowhere (TermVar leftVar) (TermVar rightVar)) leftVars rightVars datatypeCasesFormula :: CheckedDatatype -> Formula datatypeCasesFormula datatype = makeForall [witnessVar] ( isElementOf (TermVar witnessVar) datatypeCarrier `Implies` makeDisjunction disjuncts ) where usedVars = datatypeUsedVars datatype witnessVar = freshDatatypeVar usedVars "x" datatypeCarrier = datatypeCarrierTerm datatype disjuncts = datatypeCaseDisjunct witnessVar <$> NonEmpty.toList (checkedDatatypeClauses datatype) datatypeCaseDisjunct var clause = existsIfNeeded premiseVars ( makeConjunction ( premises <> [ Equals Nowhere (TermVar var) constructorTerm ] ) ) where premiseVars = datatypeClausePremiseVars clause premises = datatypeClausePremiseFormulas clause constructorTerm = datatypeClauseTerm clause (TermVar <$> checkedDatatypeClauseConstructorArgs clause) datatypeInductFormula :: CheckedDatatype -> Formula datatypeInductFormula datatype = makeForall [subsetVar] (impliesFrom closureAssumptions conclusion) where usedVars = datatypeUsedVars datatype subsetVar = freshDatatypeVar usedVars "S" witnessVar = freshDatatypeVar (Set.insert subsetVar usedVars) "x" datatypeCarrier = datatypeCarrierTerm datatype conclusion = makeForall [witnessVar] ( isElementOf (TermVar witnessVar) datatypeCarrier `Implies` isElementOf (TermVar witnessVar) (TermVar subsetVar) ) closureAssumptions = datatypeInductionClosure subsetVar <$> NonEmpty.toList (checkedDatatypeClauses datatype) datatypeInductionClosure :: VarSymbol -> CheckedDatatypeClause -> Formula datatypeInductionClosure subsetVar clause = forallIfNeeded premiseVars (impliesFrom inductionPremises conclusion) where premiseVars = datatypeClausePremiseVars clause inductionPremises = datatypeInductionPremise subsetVar <$> checkedDatatypeClausePremises clause conclusion = isElementOf ( datatypeClauseTerm clause (TermVar <$> checkedDatatypeClauseConstructorArgs clause) ) (TermVar subsetVar) datatypeInductionPremise :: VarSymbol -> CheckedDatatypePremise -> Formula datatypeInductionPremise subsetVar = \case RecursiveDatatypePremise var _canonicalDomain -> isElementOf (TermVar var) (TermVar subsetVar) NonRecursiveDatatypePremise var domain -> isElementOf (TermVar var) domain datatypeCarrierTerm :: CheckedDatatype -> Expr datatypeCarrierTerm datatype = TermSymbol Nowhere (SymbolMixfix (checkedDatatypeHead datatype)) [] datatypeClauseResultFormula :: CheckedDatatype -> CheckedDatatypeClause -> Formula datatypeClauseResultFormula datatype clause = isElementOf ( datatypeClauseTerm clause (TermVar <$> checkedDatatypeClauseConstructorArgs clause) ) (datatypeCarrierTerm datatype) datatypeClauseTerm :: CheckedDatatypeClause -> [Expr] -> Expr datatypeClauseTerm clause args = TermSymbol Nowhere (SymbolMixfix (checkedDatatypeClauseConstructor clause)) args datatypeClausePremiseVars :: CheckedDatatypeClause -> [VarSymbol] datatypeClausePremiseVars = fmap datatypePremiseVar . checkedDatatypeClausePremises datatypeClausePremiseFormulas :: CheckedDatatypeClause -> [Formula] datatypeClausePremiseFormulas = fmap datatypePremiseFormula . checkedDatatypeClausePremises datatypePremiseFormula :: CheckedDatatypePremise -> Formula datatypePremiseFormula = \case RecursiveDatatypePremise var canonicalDomain -> isElementOf (TermVar var) canonicalDomain NonRecursiveDatatypePremise var domain -> isElementOf (TermVar var) domain datatypePremiseVar :: CheckedDatatypePremise -> VarSymbol datatypePremiseVar = \case RecursiveDatatypePremise var _canonicalDomain -> var NonRecursiveDatatypePremise var _domain -> var datatypeClauseSymbolText :: CheckedDatatypeClause -> Text datatypeClauseSymbolText = functionSymbolText . checkedDatatypeClauseConstructor functionSymbolText :: FunctionSymbol -> Text functionSymbolText symbol = case mixfixMarker symbol of Marker name -> name datatypeUsedVars :: CheckedDatatype -> Set VarSymbol datatypeUsedVars datatype = Set.fromList [ var | clause <- NonEmpty.toList (checkedDatatypeClauses datatype) , var <- datatypeClausePremiseVars clause ] duplicateDatatypeConstructorSymbols :: [FunctionSymbol] -> Set FunctionSymbol duplicateDatatypeConstructorSymbols = snd . foldl' step (Set.empty, Set.empty) where step (seen, duplicates) symbol | symbol `Set.member` seen = (seen, Set.insert symbol duplicates) | otherwise = (Set.insert symbol seen, duplicates) duplicateVars :: [VarSymbol] -> Set VarSymbol duplicateVars = snd . foldl' step (Set.empty, Set.empty) where step (seen, duplicates) var | var `Set.member` seen = (seen, Set.insert var duplicates) | otherwise = (Set.insert var seen, duplicates) formatDatatypeConstructors :: Set FunctionSymbol -> Text formatDatatypeConstructors symbols = Text.intercalate ", " (functionSymbolText <$> Set.toList symbols) formatVars :: Set VarSymbol -> Text formatVars vars = Text.intercalate ", " (formatVar <$> Set.toList vars) formatVar :: VarSymbol -> Text formatVar = \case NamedVar name -> name FreshVar index -> "_" <> Text.pack (show index) renameDatatypeVars :: Set VarSymbol -> [VarSymbol] -> [VarSymbol] renameDatatypeVars _used [] = [] renameDatatypeVars used (var:vars) = let renamed = freshDatatypeLikeVar used var in renamed : renameDatatypeVars (Set.insert renamed used) vars freshDatatypeLikeVar :: Set VarSymbol -> VarSymbol -> VarSymbol freshDatatypeLikeVar used = \case NamedVar name -> freshDatatypeVar used (name <> "_rhs") FreshVar index -> freshDatatypeVar used ("_" <> Text.pack (show index) <> "_rhs") freshDatatypeVar :: Set VarSymbol -> Text -> VarSymbol freshDatatypeVar used base = List.head [ NamedVar candidate | candidate <- base : [ base <> Text.pack (show index) | index <- [(1 :: Int) ..] ] , NamedVar candidate `Set.notMember` used ] forallIfNeeded :: [VarSymbol] -> Formula -> Formula forallIfNeeded [] formula = formula forallIfNeeded vars formula = makeForall vars formula existsIfNeeded :: [VarSymbol] -> Formula -> Formula existsIfNeeded [] formula = formula existsIfNeeded vars formula = makeExists vars formula impliesFrom :: [Formula] -> Formula -> Formula impliesFrom [] conclusion = conclusion impliesFrom premises conclusion = makeConjunction premises `Implies` conclusion unorderedPairs :: [a] -> [(a, a)] unorderedPairs = \case [] -> [] value:rest -> [(value, other) | other <- rest] <> unorderedPairs rest