packages feed

proarrow-0.3.0.0: src/Proarrow/Tools/SMC/Internal/Context.hs

{-# LANGUAGE AllowAmbiguousTypes #-}

-- | Internal module of "Proarrow.Tools.SMC": the contexts of terms, and the operations on them that
-- reorder, copy and discard variables. It exports everything, also what the public module keeps
-- hidden.
module Proarrow.Tools.SMC.Internal.Context where

import Data.Kind (Constraint, Type)
import GHC.TypeNats (CmpNat, Nat)
import Proarrow.Category.Monoidal
  ( Monoidal (..)
  , SymMonoidal (..)
  , associator'
  , associatorInv'
  , rightUnitorWith
  , swapInner
  )
import Proarrow.Category.Monoidal qualified as M
import Proarrow.Core (CategoryOf (..), Promonad (..), obj)
import Proarrow.Monoid (CocommutativeComonoid, Comonoid (..))
import Proarrow.Object (Obj)
import Prelude (Ordering (..), type (~))

import Proarrow.Tools.SMC.Internal.Syntax

-- | A context: the variables a term uses, each with its id and type, by descending id.
type Ctx :: Type -> Type
type Ctx k = [(Nat, SYN k)]

-- | The type standing in for a context: the tensor of its variables' types, with the most
-- recently bound variable on the right. A single variable is just its type, so a variable is the
-- identity. The cost is that @Mul ('(n, a) ': g)@ only reduces once @g@ is known to be empty or
-- not, which 'ctxCase' tells.
type Mul :: forall {k}. Ctx k -> SYN k
type family Mul g where
  Mul '[] = I
  Mul '[ '(n, a)] = a
  Mul ('(n, a) ': g) = Mul g :** a

-- | A context that is known to be empty or not, all the way down.
type KnownCtx :: forall {k}. Ctx k -> Constraint
class KnownCtx (g :: Ctx k) where
  -- | Case analysis on the context, which is what lets @'Mul' ('(n, a) ': g)@ reduce.
  ctxCase :: ((g ~ '[]) => r) -> (forall n a g'. (g ~ ('(n, a) ': g'), KnownObj a, KnownCtx g') => r) -> r

  -- | The tensor of a context is an object. A method rather than a function over 'ctxCase', so that
  -- at a known context it is not recursive and can be inlined.
  withCtxOb :: (Monoidal k) => ((Ob (Interp (Mul g))) => r) -> r

instance KnownCtx ('[] :: Ctx k) where
  {-# INLINE ctxCase #-}
  {-# INLINE withCtxOb #-}
  ctxCase e _ = e
  withCtxOb r = r

instance (KnownObj a, KnownCtx g) => KnownCtx ('(n, a) ': g) where
  {-# INLINE ctxCase #-}
  {-# INLINE withCtxOb #-}
  ctxCase _ c = c
  withCtxOb r = ctxCase @g (withSynOb @a r) (withCtxOb @g (withSynOb @a (withOb2 @_ @(Interp (Mul g)) @(Interp a) r)))

-- | The identity on the tensor of a context.
{-# INLINE ctxOb #-}
ctxOb :: forall {k} (g :: Ctx k). (Monoidal k, KnownCtx g) => Obj (Interp (Mul g))
ctxOb = withCtxOb @g (obj @(Interp (Mul g)))

-- | A new variable on the right of a context: a unitor if the context was empty, and nothing
-- otherwise.
{-# INLINE snoc #-}
snoc
  :: forall {k} n (a :: SYN k) g
   . (Monoidal k, KnownObj a, KnownCtx g) => Interp (Mul g) ** Interp a ~> Interp (Mul ('(n, a) ': g))
snoc = ctxCase @g (withSynOb @a leftUnitor) (ctxOb @('(n, a) ': g))

-- | The context of two terms used side by side.
type Union :: forall {k}. Ctx k -> Ctx k -> Ctx k
type family Union g1 g2 where
  Union '[] g2 = g2
  Union g1 '[] = g1
  Union ('(n, a) ': g1) ('(m, b) ': g2) = UnionBy (CmpNat n m) ('(n, a) ': g1) ('(m, b) ': g2)

type UnionBy :: forall {k}. Ordering -> Ctx k -> Ctx k -> Ctx k
type family UnionBy o g1 g2 where
  UnionBy GT (x ': g1) g2 = x ': Union g1 g2
  UnionBy LT g1 (y ': g2) = y ': Union g1 g2
  UnionBy EQ (x ': g1) (y ': g2) = x ': Union g1 g2

-- | The newest variable split off the tensor of a context, the inverse of 'snoc'.
{-# INLINE unsnoc #-}
unsnoc
  :: forall {k} n (a :: SYN k) g
   . (Monoidal k, KnownObj a, KnownCtx g) => Interp (Mul ('(n, a) ': g)) ~> Interp (Mul g) ** Interp a
unsnoc = ctxCase @g (withSynOb @a leftUnitorInv) (ctxOb @('(n, a) ': g))

-- | Split the tensor of the union of two contexts into the tensors of the two contexts. This is
-- where the wires are reordered, the only place 'swap' is used, and where a variable that both
-- contexts have is copied.
type Merge :: forall {k}. Ctx k -> Ctx k -> Constraint
class (KnownCtx g1, KnownCtx g2) => Merge (g1 :: Ctx k) g2 where
  merge :: Interp (Mul (Union g1 g2)) ~> Interp (Mul g1) ** Interp (Mul g2)

instance (Monoidal k, KnownCtx g2) => Merge ('[] :: Ctx k) g2 where
  {-# INLINE merge #-}
  merge = withCtxOb @g2 leftUnitorInv

instance (Monoidal k, KnownCtx ('(n, a) ': g1)) => Merge ('(n, a) ': g1 :: Ctx k) '[] where
  {-# INLINE merge #-}
  merge = withCtxOb @('(n, a) ': g1) rightUnitorInv

instance
  ( Monoidal k
  , KnownObj a
  , KnownObj b
  , KnownCtx g1
  , KnownCtx g2
  , MergeBy (CmpNat n m) ('(n, a) ': g1 :: Ctx k) ('(m, b) ': g2)
  )
  => Merge ('(n, a) ': g1 :: Ctx k) ('(m, b) ': g2)
  where
  {-# INLINE merge #-}
  merge = mergeBy @(CmpNat n m) @('(n, a) ': g1) @('(m, b) ': g2)

-- | 'merge' for two non-empty contexts, by which of the two has the larger head id.
type MergeBy :: forall {k}. Ordering -> Ctx k -> Ctx k -> Constraint
class (KnownCtx g1, KnownCtx g2) => MergeBy o (g1 :: Ctx k) g2 where
  mergeBy :: Interp (Mul (UnionBy o g1 g2)) ~> Interp (Mul g1) ** Interp (Mul g2)

-- The union of two non-empty contexts is not empty, so its tensor splits off the newest variable
-- as is, which the equality says for GHC. If the newest variable is alone on its side, merging is
-- one swap or nothing.
instance
  ( SymMonoidal k
  , Merge g1 ('(m, b) ': g2)
  , KnownObj (a :: SYN k)
  , KnownObj b
  , Mul ('(n, a) ': Union g1 ('(m, b) ': g2)) ~ (Mul (Union g1 ('(m, b) ': g2)) :** a)
  )
  => MergeBy GT ('(n, a) ': g1) ('(m, b) ': g2)
  where
  {-# INLINE mergeBy #-}
  mergeBy =
    withCtxOb @('(m, b) ': g2)
      ( withSynOb @a
          ( ctxCase @g1
              (swap @k @(Interp (Mul ('(m, b) ': g2))) @(Interp a))
              ( associatorInv' (ctxOb @g1) (synOb @a) (ctxOb @('(m, b) ': g2))
                  . (ctxOb @g1 M.** swap @k @(Interp (Mul ('(m, b) ': g2))) @(Interp a))
                  . associator' (ctxOb @g1) (ctxOb @('(m, b) ': g2)) (synOb @a)
                  . (merge @g1 @('(m, b) ': g2) M.** synOb @a)
              )
          )
      )

instance
  ( Monoidal k
  , Merge ('(n, a) ': g1) g2
  , KnownObj a
  , KnownObj (b :: SYN k)
  , Mul ('(m, b) ': Union ('(n, a) ': g1) g2) ~ (Mul (Union ('(n, a) ': g1) g2) :** b)
  )
  => MergeBy LT ('(n, a) ': g1) ('(m, b) ': g2)
  where
  {-# INLINE mergeBy #-}
  mergeBy =
    ctxCase @g2
      (ctxOb @('(m, b) ': '(n, a) ': g1))
      ( associator' (ctxOb @('(n, a) ': g1)) (ctxOb @g2) (synOb @b)
          . (merge @('(n, a) ': g1) @g2 M.** synOb @b)
      )

-- | A variable both contexts have is copied, one copy for each side.
instance
  ( SymMonoidal k
  , CocommutativeComonoid (Interp a)
  , KnownObj (a :: SYN k)
  , a ~ b
  , Merge g1 g2
  , KnownCtx (Union g1 g2)
  )
  => MergeBy EQ ('(n, a) ': g1) ('(m, b) ': g2)
  where
  {-# INLINE mergeBy #-}
  mergeBy = ctxCase @g1 (ctxCase @g2 (comult @(Interp a)) withRest) withRest
    where
      -- When the variable is all both sides have, this is just the copy.
      withRest :: Interp (Mul ('(n, a) ': Union g1 g2)) ~> Interp (Mul ('(n, a) ': g1)) ** Interp (Mul ('(m, b) ': g2))
      withRest =
        withCtxOb @g1
          ( withCtxOb @g2
              ( withSynOb @a
                  ( (snoc @n @a @g1 M.** snoc @m @a @g2)
                      . swapInner @(Interp (Mul g1)) @(Interp (Mul g2)) @(Interp a) @(Interp a)
                      . (merge @g1 @g2 M.** comult @(Interp a))
                      . unsnoc @n @a @(Union g1 g2)
                  )
              )
          )

-- | The inverse of 'merge' for two contexts without a variable in common, which a @rec@ block
-- uses to put the wires it feeds back together with the others.
type Unmerge :: forall {k}. Ctx k -> Ctx k -> Constraint
class (KnownCtx g1, KnownCtx g2) => Unmerge (g1 :: Ctx k) g2 where
  unmerge :: Interp (Mul g1) ** Interp (Mul g2) ~> Interp (Mul (Union g1 g2))

instance (Monoidal k, KnownCtx g2) => Unmerge ('[] :: Ctx k) g2 where
  {-# INLINE unmerge #-}
  unmerge = withCtxOb @g2 leftUnitor

instance (Monoidal k, KnownCtx ('(n, a) ': g1)) => Unmerge ('(n, a) ': g1 :: Ctx k) '[] where
  {-# INLINE unmerge #-}
  unmerge = withCtxOb @('(n, a) ': g1) rightUnitor

instance
  ( Monoidal k
  , KnownObj a
  , KnownObj b
  , KnownCtx g1
  , KnownCtx g2
  , UnmergeBy (CmpNat n m) ('(n, a) ': g1 :: Ctx k) ('(m, b) ': g2)
  )
  => Unmerge ('(n, a) ': g1 :: Ctx k) ('(m, b) ': g2)
  where
  {-# INLINE unmerge #-}
  unmerge = unmergeBy @(CmpNat n m) @('(n, a) ': g1) @('(m, b) ': g2)

-- | 'unmerge' for two non-empty contexts, by which of the two has the larger head id.
type UnmergeBy :: forall {k}. Ordering -> Ctx k -> Ctx k -> Constraint
class (KnownCtx g1, KnownCtx g2) => UnmergeBy o (g1 :: Ctx k) g2 where
  unmergeBy :: Interp (Mul g1) ** Interp (Mul g2) ~> Interp (Mul (UnionBy o g1 g2))

instance
  ( SymMonoidal k
  , Unmerge g1 ('(m, b) ': g2)
  , KnownObj (a :: SYN k)
  , KnownObj b
  , Mul ('(n, a) ': Union g1 ('(m, b) ': g2)) ~ (Mul (Union g1 ('(m, b) ': g2)) :** a)
  )
  => UnmergeBy GT ('(n, a) ': g1) ('(m, b) ': g2)
  where
  {-# INLINE unmergeBy #-}
  unmergeBy =
    withCtxOb @('(m, b) ': g2)
      ( withSynOb @a
          ( ctxCase @g1
              (swap @k @(Interp a) @(Interp (Mul ('(m, b) ': g2))))
              ( (unmerge @g1 @('(m, b) ': g2) M.** synOb @a)
                  . associatorInv' (ctxOb @g1) (ctxOb @('(m, b) ': g2)) (synOb @a)
                  . (ctxOb @g1 M.** swap @k @(Interp a) @(Interp (Mul ('(m, b) ': g2))))
                  . associator' (ctxOb @g1) (synOb @a) (ctxOb @('(m, b) ': g2))
              )
          )
      )

instance
  ( Monoidal k
  , Unmerge ('(n, a) ': g1) g2
  , KnownObj a
  , KnownObj (b :: SYN k)
  , Mul ('(m, b) ': Union ('(n, a) ': g1) g2) ~ (Mul (Union ('(n, a) ': g1) g2) :** b)
  )
  => UnmergeBy LT ('(n, a) ': g1) ('(m, b) ': g2)
  where
  {-# INLINE unmergeBy #-}
  unmergeBy =
    ctxCase @g2
      (ctxOb @('(m, b) ': '(n, a) ': g1))
      ( (unmerge @('(n, a) ': g1) @g2 M.** synOb @b)
          . associatorInv' (ctxOb @('(n, a) ': g1)) (ctxOb @g2) (synOb @b)
      )

-- | Keep the variables of @g@ that @h@ has and discard the others, for an alternative of the
-- additives that does not use all the variables of the other. Every variable of @h@ must be in
-- @g@.
type Thin :: forall {k}. Ctx k -> Ctx k -> Constraint
class (KnownCtx g, KnownCtx h) => Thin (g :: Ctx k) h where
  thin :: Interp (Mul g) ~> Interp (Mul h)

instance (Monoidal k) => Thin ('[] :: Ctx k) '[] where
  {-# INLINE thin #-}
  thin = obj @(Unit :: k)

instance (Monoidal k, Comonoid (Interp a), KnownObj (a :: SYN k), Thin g '[]) => Thin ('(n, a) ': g) '[] where
  {-# INLINE thin #-}
  thin = leftUnitor @k @Unit . (thin @g @'[] M.** counit @(Interp a)) . unsnoc @n @a @g

instance
  (KnownObj a, KnownObj b, KnownCtx g, KnownCtx h, ThinBy (CmpNat n m) ('(n, a) ': g) ('(m, b) ': h))
  => Thin ('(n, a) ': g :: Ctx k) ('(m, b) ': h)
  where
  {-# INLINE thin #-}
  thin = thinBy @(CmpNat n m) @('(n, a) ': g) @('(m, b) ': h)

-- | 'thin' for two non-empty contexts, by which of the two has the larger head id.
type ThinBy :: forall {k}. Ordering -> Ctx k -> Ctx k -> Constraint
class (KnownCtx g, KnownCtx h) => ThinBy o (g :: Ctx k) h where
  thinBy :: Interp (Mul g) ~> Interp (Mul h)

instance (Monoidal k, KnownObj (a :: SYN k), a ~ b, Thin g h) => ThinBy EQ ('(n, a) ': g) ('(m, b) ': h) where
  {-# INLINE thinBy #-}
  thinBy = withSynOb @a (snoc @m @a @h . (thin @g @h M.** synOb @a) . unsnoc @n @a @g)

instance
  (Monoidal k, Comonoid (Interp a), KnownObj (a :: SYN k), KnownObj b, Thin g ('(m, b) ': h))
  => ThinBy GT ('(n, a) ': g) ('(m, b) ': h)
  where
  {-# INLINE thinBy #-}
  thinBy =
    withCtxOb @('(m, b) ': h)
      (rightUnitor . (thin @g @('(m, b) ': h) M.** counit @(Interp a)) . unsnoc @n @a @g)

-- | A binder's variable @n@ on the right of @r@, what the body uses besides it, given the body's
-- context @g@: 'snoc' if the body uses the variable, and its counit otherwise. The binder's
-- variable is the newest in scope, so it can only be at the head of @g@. The functional
-- dependency, rather than a type family, gives @r@, so that a binder costs one comparison of ids.
type BindVar :: forall {k}. Nat -> SYN k -> Ctx k -> Ctx k -> Constraint
class (KnownCtx r) => BindVar n (a :: SYN k) g r | n g -> r where
  bindVar :: Interp (Mul r) ** Interp a ~> Interp (Mul g)

instance (Monoidal k, Comonoid (Interp a), KnownObj (a :: SYN k)) => BindVar n a '[] '[] where
  {-# INLINE bindVar #-}
  bindVar = rightUnitorWith @Unit (counit @(Interp a))

instance (KnownCtx r, BindVarBy (CmpNat n m) n a ('(m, b) ': g) r) => BindVar n (a :: SYN k) ('(m, b) ': g) r where
  {-# INLINE bindVar #-}
  bindVar = bindVarBy @(CmpNat n m) @n @a @('(m, b) ': g) @r

-- | 'bindVar' for a non-empty context, by whether its head is the new variable.
type BindVarBy :: forall {k}. Ordering -> Nat -> SYN k -> Ctx k -> Ctx k -> Constraint
class BindVarBy o n (a :: SYN k) g r | o n g -> r where
  bindVarBy :: Interp (Mul r) ** Interp a ~> Interp (Mul g)

instance (Monoidal k, KnownObj (a :: SYN k), a ~ b, KnownCtx g) => BindVarBy EQ n a ('(m, b) ': g) g where
  {-# INLINE bindVarBy #-}
  bindVarBy = snoc @m @a @g

instance
  (Monoidal k, Comonoid (Interp a), KnownObj (a :: SYN k), KnownObj b, KnownCtx g)
  => BindVarBy GT n a ('(m, b) ': g) ('(m, b) ': g)
  where
  {-# INLINE bindVarBy #-}
  bindVarBy =
    withCtxOb @('(m, b) ': g)
      (rightUnitorWith @(Interp (Mul ('(m, b) ': g))) (counit @(Interp a)))

type HeadId :: forall {k}. Ctx k -> Nat
type family HeadId g where
  HeadId ('(n, a) ': g) = n

-- | The variables of @g@ whose ids are not in @h@.
type Minus :: forall {k}. Ctx k -> Ctx k -> Ctx k
type family Minus g h where
  Minus '[] h = '[]
  Minus g '[] = g
  Minus ('(n, a) ': g) ('(m, b) ': h) = MinusBy (CmpNat n m) ('(n, a) ': g) ('(m, b) ': h)

type MinusBy :: forall {k}. Ordering -> Ctx k -> Ctx k -> Ctx k
type family MinusBy o g h where
  MinusBy GT (x ': g) h = x ': Minus g h
  MinusBy EQ (x ': g) (y ': h) = Minus g h
  MinusBy LT g (y ': h) = Minus g h

-- | The variables of @g@ whose ids are in @h@.
type Inter :: forall {k}. Ctx k -> Ctx k -> Ctx k
type Inter g h = Minus g (Minus g h)