packages feed

hydra-0.12.0: src/gen-main/haskell/Hydra/Substitution.hs

-- | Variable substitution in type and term expressions.

module Hydra.Substitution where

import qualified Hydra.Core as Core
import qualified Hydra.Lib.Lists as Lists
import qualified Hydra.Lib.Logic as Logic
import qualified Hydra.Lib.Maps as Maps
import qualified Hydra.Lib.Optionals as Optionals
import qualified Hydra.Lib.Sets as Sets
import qualified Hydra.Rewriting as Rewriting
import qualified Hydra.Typing as Typing
import Prelude hiding  (Enum, Ordering, fail, map, pure, sum)
import qualified Data.Int as I
import qualified Data.List as L
import qualified Data.Map as M
import qualified Data.Set as S

composeTypeSubst :: (Typing.TypeSubst -> Typing.TypeSubst -> Typing.TypeSubst)
composeTypeSubst s1 s2 =  
  let isExtra = (\k -> \v -> Optionals.isNothing (Maps.lookup k (Typing.unTypeSubst s1))) 
      withExtra = (Maps.filterWithKey isExtra (Typing.unTypeSubst s2))
  in (Typing.TypeSubst (Maps.union withExtra (Maps.map (substInType s2) (Typing.unTypeSubst s1))))

composeTypeSubstList :: ([Typing.TypeSubst] -> Typing.TypeSubst)
composeTypeSubstList = (Lists.foldl composeTypeSubst idTypeSubst)

idTypeSubst :: Typing.TypeSubst
idTypeSubst = (Typing.TypeSubst Maps.empty)

singletonTypeSubst :: (Core.Name -> Core.Type -> Typing.TypeSubst)
singletonTypeSubst v t = (Typing.TypeSubst (Maps.singleton v t))

substituteInConstraint :: (Typing.TypeSubst -> Typing.TypeConstraint -> Typing.TypeConstraint)
substituteInConstraint subst c = Typing.TypeConstraint {
  Typing.typeConstraintLeft = (substInType subst (Typing.typeConstraintLeft c)),
  Typing.typeConstraintRight = (substInType subst (Typing.typeConstraintRight c)),
  Typing.typeConstraintComment = (Typing.typeConstraintComment c)}

substituteInConstraints :: (Typing.TypeSubst -> [Typing.TypeConstraint] -> [Typing.TypeConstraint])
substituteInConstraints subst cs = (Lists.map (substituteInConstraint subst) cs)

substInContext :: (Typing.TypeSubst -> Typing.InferenceContext -> Typing.InferenceContext)
substInContext subst cx = Typing.InferenceContext {
  Typing.inferenceContextSchemaTypes = (Typing.inferenceContextSchemaTypes cx),
  Typing.inferenceContextPrimitiveTypes = (Typing.inferenceContextPrimitiveTypes cx),
  Typing.inferenceContextDataTypes = (Maps.map (substInTypeScheme subst) (Typing.inferenceContextDataTypes cx)),
  Typing.inferenceContextDebug = (Typing.inferenceContextDebug cx)}

substituteInTerm :: (Typing.TermSubst -> Core.Term -> Core.Term)
substituteInTerm subst =  
  let s = (Typing.unTermSubst subst) 
      rewrite = (\recurse -> \term ->  
              let withLambda = (\l ->  
                      let v = (Core.lambdaParameter l) 
                          subst2 = (Typing.TermSubst (Maps.remove v s))
                      in (Core.TermFunction (Core.FunctionLambda (Core.Lambda {
                        Core.lambdaParameter = v,
                        Core.lambdaDomain = (Core.lambdaDomain l),
                        Core.lambdaBody = (substituteInTerm subst2 (Core.lambdaBody l))})))) 
                  withLet = (\lt ->  
                          let bindings = (Core.letBindings lt) 
                              names = (Sets.fromList (Lists.map Core.bindingName bindings))
                              subst2 = (Typing.TermSubst (Maps.filterWithKey (\k -> \v -> Logic.not (Sets.member k names)) s))
                              rewriteBinding = (\b -> Core.Binding {
                                      Core.bindingName = (Core.bindingName b),
                                      Core.bindingTerm = (substituteInTerm subst2 (Core.bindingTerm b)),
                                      Core.bindingType = (Core.bindingType b)})
                          in (Core.TermLet (Core.Let {
                            Core.letBindings = (Lists.map rewriteBinding bindings),
                            Core.letEnvironment = (substituteInTerm subst2 (Core.letEnvironment lt))})))
              in ((\x -> case x of
                Core.TermFunction v1 -> ((\x -> case x of
                  Core.FunctionLambda v2 -> (withLambda v2)
                  _ -> (recurse term)) v1)
                Core.TermLet v1 -> (withLet v1)
                Core.TermVariable v1 -> (Optionals.maybe (recurse term) (\sterm -> sterm) (Maps.lookup v1 s))
                _ -> (recurse term)) term))
  in (Rewriting.rewriteTerm rewrite)

substInType :: (Typing.TypeSubst -> Core.Type -> Core.Type)
substInType subst =  
  let rewrite = (\recurse -> \typ -> (\x -> case x of
          Core.TypeForall v1 -> (Optionals.maybe (recurse typ) (\styp -> Core.TypeForall (Core.ForallType {
            Core.forallTypeParameter = (Core.forallTypeParameter v1),
            Core.forallTypeBody = (substInType (removeVar (Core.forallTypeParameter v1)) (Core.forallTypeBody v1))})) (Maps.lookup (Core.forallTypeParameter v1) (Typing.unTypeSubst subst)))
          Core.TypeVariable v1 -> (Optionals.maybe typ (\styp -> styp) (Maps.lookup v1 (Typing.unTypeSubst subst)))
          _ -> (recurse typ)) typ) 
      removeVar = (\v -> Typing.TypeSubst (Maps.remove v (Typing.unTypeSubst subst)))
  in (Rewriting.rewriteType rewrite)

substInTypeScheme :: (Typing.TypeSubst -> Core.TypeScheme -> Core.TypeScheme)
substInTypeScheme subst ts = Core.TypeScheme {
  Core.typeSchemeVariables = (Core.typeSchemeVariables ts),
  Core.typeSchemeType = (substInType subst (Core.typeSchemeType ts))}

substTypesInTerm :: (Typing.TypeSubst -> Core.Term -> Core.Term)
substTypesInTerm subst =  
  let rewrite = (\recurse -> \term ->  
          let forElimination = (\elm -> (\x -> case x of
                  Core.EliminationProduct v1 -> (forTupleProjection v1)
                  _ -> (recurse term)) elm) 
              forFunction = (\f -> (\x -> case x of
                      Core.FunctionElimination v1 -> (forElimination v1)
                      Core.FunctionLambda v1 -> (forLambda v1)
                      _ -> (recurse term)) f)
              forLambda = (\l -> recurse (Core.TermFunction (Core.FunctionLambda (Core.Lambda {
                      Core.lambdaParameter = (Core.lambdaParameter l),
                      Core.lambdaDomain = (Optionals.map (substInType subst) (Core.lambdaDomain l)),
                      Core.lambdaBody = (Core.lambdaBody l)}))))
              forLet = (\l ->  
                      let rewriteBinding = (\b -> Core.Binding {
                              Core.bindingName = (Core.bindingName b),
                              Core.bindingTerm = (Core.bindingTerm b),
                              Core.bindingType = (Optionals.map (substInTypeScheme subst) (Core.bindingType b))})
                      in (recurse (Core.TermLet (Core.Let {
                        Core.letBindings = (Lists.map rewriteBinding (Core.letBindings l)),
                        Core.letEnvironment = (Core.letEnvironment l)}))))
              forTupleProjection = (\tp -> recurse (Core.TermFunction (Core.FunctionElimination (Core.EliminationProduct (Core.TupleProjection {
                      Core.tupleProjectionArity = (Core.tupleProjectionArity tp),
                      Core.tupleProjectionIndex = (Core.tupleProjectionIndex tp),
                      Core.tupleProjectionDomain = (Optionals.map (\types -> Lists.map (substInType subst) types) (Core.tupleProjectionDomain tp))})))))
              forTypeLambda = (\ta ->  
                      let param = (Core.typeLambdaParameter ta) 
                          subst2 = (Typing.TypeSubst (Maps.remove param (Typing.unTypeSubst subst)))
                      in (Core.TermTypeLambda (Core.TypeLambda {
                        Core.typeLambdaParameter = param,
                        Core.typeLambdaBody = (substTypesInTerm subst2 (Core.typeLambdaBody ta))})))
          in ((\x -> case x of
            Core.TermFunction v1 -> (forFunction v1)
            Core.TermLet v1 -> (forLet v1)
            Core.TermTypeLambda v1 -> (forTypeLambda v1)
            Core.TermTypeApplication v1 -> (recurse (Core.TermTypeApplication (Core.TypedTerm {
              Core.typedTermTerm = (Core.typedTermTerm v1),
              Core.typedTermType = (substInType subst (Core.typedTermType v1))})))
            _ -> (recurse term)) term))
  in (Rewriting.rewriteTerm rewrite)