0 | module Syntax.StringDiagram.Arrow
  1 |
  2 | import public Control.Monad.State
  3 | import public Control.Applicative.Const
  4 | import public Control.Arrow
  5 | import public Control.Category.Core
  6 | import public Language.Reflection
  7 | import public Data.Wrap0
  8 |
  9 | %default total
 10 | %language ElabReflection
 11 |
 12 | export infix 0 -<
 13 | export prefix 0 =<
 14 |
 15 |
 16 | record ArrowVar where
 17 |   constructor MkAVar
 18 |   name : String
 19 |   type : TTImp
 20 |
 21 |
 22 | zipFilter : List a -> List Bool -> List a
 23 | zipFilter = catMaybes .: zipWith (\x,b => guard b $> x)
 24 |
 25 |
 26 | bindVars : TTImp -> List String
 27 | bindVars = (nub . runConst .) $ mapATTImp' $ \case
 28 |   IBindVar _ (UN $ Basic var) => const $ MkConst [var]
 29 |   IAs _ _ _ (UN $ Basic var) _ => \(MkConst vs) => MkConst (var :: vs)
 30 |   _ => id
 31 |
 32 | isBoundVar : TTImp -> String -> Bool
 33 | isBoundVar t var =
 34 |   let _ = Lazy.Monoid.Any
 35 |   in runConst {a=Lazy Bool} $ flip mapATTImp' t $ \case
 36 |     IBindVar _ (UN $ Basic var') => const $ MkConst $ delay (var == var')
 37 |     IAs _ _ _ (UN $ Basic var') _ => \(MkConst b) => MkConst (var == var' || b)
 38 |     _ => id
 39 |
 40 | usedInExpr : TTImp -> String -> Bool
 41 | usedInExpr t var =
 42 |   let _ = Lazy.Monoid.Any
 43 |   in runConst {a=Lazy Bool} $ flip mapATTImp' t $ \case
 44 |     -- If a hole is found, always include every possible variable
 45 |     IHole _ _ => const $ MkConst $ delay True
 46 |     IVar _ (UN $ Basic var') => const $ MkConst $ delay (var == var')
 47 |     _ => id
 48 |
 49 | usedInDiagram : TTImp -> String -> Bool
 50 | usedInDiagram (ILam _ _ _ (Just $ UN $ Basic v) _ rest) var =
 51 |   v /= var && usedInDiagram rest var
 52 | usedInDiagram (ILam _ _ _ _ _ rest) var = usedInDiagram rest var
 53 | usedInDiagram (ILet _ _ _ (UN $ Basic v) _ exp rest) var =
 54 |   usedInExpr exp var || (v /= var && usedInDiagram rest var)
 55 | usedInDiagram (ILet _ _ _ _ _ exp rest) var =
 56 |   usedInExpr exp var || usedInDiagram rest var
 57 | usedInDiagram (ICase _ _ exp _ clauses) var =
 58 |   usedInExpr exp var || any (\case
 59 |     PatClause _ lhs rhs =>
 60 |       not (isBoundVar lhs var) &&
 61 |         usedInDiagram (assert_smaller clauses rhs) var
 62 |     _ => False
 63 |     ) clauses
 64 | usedInDiagram `((~(mor) -< ~(inp)) >>= ~(rest)) var =
 65 |   usedInExpr inp var || usedInDiagram rest var
 66 | usedInDiagram `((~(mor) -< ~(inp)) >> ~(rest)) var =
 67 |   usedInExpr inp var || usedInDiagram rest var
 68 | usedInDiagram `(=< ~(out)) var = usedInExpr out var
 69 | usedInDiagram _ _ = False
 70 |
 71 |
 72 | export
 73 | arrowImpl : TTImp -> Elab TTImp
 74 | arrowImpl t = do
 75 |   ts <- arrowImplLam [<] [] t
 76 |   pure $ composeImp $ optimize ts
 77 |   where
 78 |     optimize : SnocList TTImp -> SnocList TTImp
 79 |     optimize [<] = [<]
 80 |     optimize [<t] = [<t]
 81 |     optimize (ts :< `(Control.Arrow.arrow {a = ~a, b = ~_} ~f) :< `(Control.Arrow.arrow {a = ~b, b = ~c} ~g)) =
 82 |       assert_total $ optimize (ts :< `(Control.Arrow.arrow {a = ~a, b = ~c} (Prelude.(.) {a = ~a, b = ~b, c = ~c} ~g ~f)))
 83 |     optimize (ts :< t :< t') = optimize (ts :< t) :< t'
 84 |
 85 |     composeImp : SnocList TTImp -> TTImp
 86 |     composeImp [<] = `(Control.Category.Core.id)
 87 |     composeImp [<t] = t
 88 |     composeImp (ts :< t) =
 89 |       `(Control.Category.Core.(.)
 90 |         ~t ~(composeImp ts))
 91 |
 92 |     pairTy : List ArrowVar -> TTImp
 93 |     pairTy [] = `(Builtin.Unit)
 94 |     pairTy [MkAVar _ ty] = ty
 95 |     pairTy (MkAVar _ ty :: vars) = `(Builtin.Pair ~ty ~(pairTy vars))
 96 |
 97 |     pair : List String -> TTImp
 98 |     pair [] = `(Builtin.MkUnit)
 99 |     pair [n] = IVar EmptyFC $ UN $ Basic n
100 |     pair (n :: ns) = `(Builtin.MkPair ~(IVar EmptyFC $ UN $ Basic n) ~(pair ns))
101 |
102 |     stringBind : List (ArrowVar, Bool) -> TTImp
103 |     stringBind [] = `(Builtin.MkUnit)
104 |     stringBind [(v,b)] = if b then IBindVar EmptyFC (UN $ Basic v.name) else Implicit EmptyFC True
105 |     stringBind ((v,b) :: strings) = `(Builtin.MkPair ~(if b then IBindVar EmptyFC (UN $ Basic v.name) else Implicit EmptyFC True) ~(stringBind strings))
106 |
107 |     stringLam : List (ArrowVar, Bool) -> TTImp -> TTImp
108 |     stringLam [] t = `(\_ => ~t)
109 |     stringLam [(_,False)] t = `(\_ => ~t)
110 |     stringLam [(v,True)] t = ILam EmptyFC MW ExplicitArg (Just $ UN $ Basic v.name) (Implicit EmptyFC False) t
111 |     stringLam ns t = `(\ ~(stringBind ns) => ~t)
112 |
113 |     mkEither : (i,n : Nat) -> TTImp -> TTImp
114 |     mkEither Z (S (S _)) t = `(Prelude.Left ~t)
115 |     mkEither (S n) (S n') t = `(Prelude.Right ~(mkEither n n' t))
116 |     mkEither _ _ t = t
117 |
118 |     arrowImplLam : SnocList TTImp -> List ArrowVar -> TTImp -> Elab (SnocList TTImp)
119 |     arrowImpl' : SnocList TTImp -> List ArrowVar -> TTImp -> Elab (SnocList TTImp)
120 |
121 |     arrowImplLam ts [] (ILam _ MW ExplicitArg Nothing _ rest) =
122 |       arrowImpl' (ts :< `(Control.Arrow.arrow {a=_,b=_} (\_ => Builtin.MkUnit))) [] rest
123 |     arrowImplLam ts [] (ILam _ MW ExplicitArg (Just $ UN $ Basic var) ty rest) =
124 |       arrowImpl' ts [MkAVar var ty] rest
125 |     arrowImplLam ts strings (ILam _ MW ExplicitArg Nothing _ rest) =
126 |       arrowImpl' (ts :< `(Control.Arrow.arrow {a=_,b=_} Builtin.snd)) strings rest
127 |     arrowImplLam ts strings (ILam _ MW ExplicitArg (Just $ UN $ Basic var) ty rest) =
128 |       if any (\v => v.name == var) strings
129 |       then do
130 |         let used = map (\v => v.name /= var) strings
131 |         let strings' = MkAVar var ty :: zipFilter strings used
132 |         arrowImpl'
133 |           (ts :< `(Control.Arrow.arrow {a = ~(pairTy $ MkAVar var ty :: strings), b = ~(pairTy strings')}
134 |                     (\(Builtin.MkPair ~(IBindVar EmptyFC $ UN $ Basic var) ~(stringBind $ zip strings used)) =>
135 |                       ~(pair $ map name strings'))))
136 |           strings' rest
137 |       else arrowImpl' ts (MkAVar var ty :: strings) rest
138 |     arrowImplLam ts strings (ILam _ MW ExplicitArg (Just $ MN var _) ty rest) = do
139 |       arrowImpl' ts (MkAVar var ty :: strings) rest
140 |     arrowImplLam ts strings (ILam fc {}) = failAt fc "Invalid lambda"
141 |     arrowImplLam ts strings t = failAt (getFC t) "Expected lambda expression or binding pattern"
142 |
143 |     arrowImpl' ts [MkAVar _ ty] (ICase _ _ (IVar _ var@(MN {})) _ [PatClause _ lhs rhs]) = do
144 |       let names = bindVars lhs
145 |       let vars = map (`MkAVar` `(_)) names
146 |       arrowImpl'
147 |         (ts :< `(Control.Arrow.arrow {a=_,b = ~(pairTy vars)} (\ ~lhs : ~ty => ~(pair names))))
148 |         vars rhs
149 |     arrowImpl' ts (v :: strings) (ICase _ _ (IVar _ var@(MN {})) _ [PatClause _ lhs rhs]) = do
150 |       let names = bindVars lhs
151 |       let vars = map (`MkAVar` `(_)) names
152 |       rest <- genSym "vars"
153 |       arrowImpl'
154 |         (ts :< `(Control.Arrow.arrow {a = ~(pairTy $ v :: strings), b = ~(pairTy $ vars ++ strings)}
155 |           (\(Builtin.MkPair ~lhs ~(IBindVar EmptyFC rest)) =>
156 |             Builtin.MkPair ~(pair names) ~(IVar EmptyFC rest))))
157 |         (vars ++ strings) rhs
158 |     arrowImpl' ts strings (ICase _ _ exp ty clauses) = do
159 |       let tot = length clauses
160 |       (clauses', conts) <- map unzip $ evalStateT Z $ for clauses $ \case
161 |         PatClause _ lhs rhs => do
162 |           let names = bindVars lhs
163 |           let vars = map (`MkAVar` `(_)) names
164 |           let usedExp = map (usedInExpr exp . name) strings
165 |           let usedLater = map (\v => not (elem v.name names) && usedInDiagram rhs v.name) strings
166 |           let strings' = vars ++ zipFilter strings usedLater
167 |           t <- lift $ assert_total $ arrowImpl' [<] strings' rhs
168 |           i <- get
169 |           modify S
170 |           pure (PatClause EmptyFC lhs $ mkEither i tot $ pair $ map name strings', Just $ composeImp $ optimize t)
171 |         ImpossibleClause _ lhs => pure (ImpossibleClause EmptyFC lhs, Nothing)
172 |         WithClause fc {} => failAt fc "Unrecognized case pattern"
173 |       let conts' = catMaybes conts
174 |       case conts' of
175 |         [] =>
176 |           pure (ts :< `(Control.Arrow.arrow {a = ~(pairTy strings), b = _}
177 |                       ~(stringLam (map (,True) strings) (ICase EmptyFC [] exp ty clauses'))))
178 |         _ :: _ =>
179 |           pure (ts :< `(Control.Arrow.arrow {a = ~(pairTy strings), b = _}
180 |                       ~(stringLam (map (,True) strings) (ICase EmptyFC [] exp ty clauses')))
181 |                     :< foldr1 (\t,t' => `(Control.Arrow.(\|/) ~t ~t')) conts')
182 |     arrowImpl' ts strings `(~(ICase _ _ exp ty clauses) >> ~(rest)) = do
183 |       let tot = length clauses
184 |       (clauses', conts) <- map unzip $ evalStateT Z $ for clauses $ \case
185 |         PatClause _ lhs rhs => do
186 |           let names = bindVars lhs
187 |           let vars = map (`MkAVar` `(_)) names
188 |           let usedExp = map (usedInExpr exp . name) strings
189 |           let usedLater = map (\v => not (elem v.name names) && usedInDiagram rhs v.name) strings
190 |           let strings' = vars ++ zipFilter strings usedLater
191 |           t <- lift $ assert_total $ arrowImpl' [<] strings' rhs
192 |           i <- get
193 |           modify S
194 |           pure (PatClause EmptyFC lhs $ mkEither i tot $ pair $ map name strings', Just $ composeImp $ optimize t)
195 |         ImpossibleClause _ lhs => pure (ImpossibleClause EmptyFC lhs, Nothing)
196 |         WithClause fc {} => failAt fc "Unrecognized case pattern"
197 |       let conts' = catMaybes conts
198 |       case conts' of
199 |         [] =>
200 |           arrowImpl' (ts :< `(Control.Arrow.arrow {a = ~(pairTy strings), b = _}
201 |                       ~(stringLam (map (,True) strings) (ICase EmptyFC [] exp ty clauses'))))
202 |             strings rest
203 |         _ :: _ =>
204 |           arrowImpl' (ts :< `(Control.Arrow.arrow {a = ~(pairTy strings), b = _}
205 |                         ~(stringLam (map (,True) strings)
206 |                         `(Builtin.MkPair ~(ICase EmptyFC [] exp ty clauses') ~(pair $ map name strings))))
207 |                       :< `(Control.Arrow.first ~(foldr1 (\t,t' => `(Control.Arrow.(\|/) ~t ~t')) conts'))
208 |                       :< `(Control.Arrow.arrow {a=_,b=_} Builtin.snd))
209 |             strings rest
210 |     arrowImpl' ts strings `(~(ICase _ _ exp ty clauses) >>= ~(rest)) = do
211 |       let tot = length clauses
212 |       (clauses', conts) <- map unzip $ evalStateT Z $ for clauses $ \case
213 |         PatClause _ lhs rhs => do
214 |           let names = bindVars lhs
215 |           let vars = map (`MkAVar` `(_)) names
216 |           let usedExp = map (usedInExpr exp . name) strings
217 |           let usedLater = map (\v => not (elem v.name names) && usedInDiagram rhs v.name) strings
218 |           let strings' = vars ++ zipFilter strings usedLater
219 |           t <- lift $ assert_total $ arrowImpl' [<] strings' rhs
220 |           i <- get
221 |           modify S
222 |           pure (PatClause EmptyFC lhs $ mkEither i tot $ pair $ map name strings', Just $ composeImp $ optimize t)
223 |         ImpossibleClause _ lhs => pure (ImpossibleClause EmptyFC lhs, Nothing)
224 |         WithClause fc {} => failAt fc "Unrecognized case pattern"
225 |       let conts' = catMaybes conts
226 |       case conts' of
227 |         [] =>
228 |           arrowImplLam (ts :< `(Control.Arrow.arrow {a = ~(pairTy strings), b = _}
229 |                       ~(stringLam (map (,True) strings) (ICase EmptyFC [] exp ty clauses'))))
230 |             strings rest
231 |         _ :: _ =>
232 |           arrowImplLam (ts :< `(Control.Arrow.arrow {a = ~(pairTy strings), b = _}
233 |                         ~(stringLam (map (,True) strings)
234 |                         `(Builtin.MkPair ~(ICase EmptyFC [] exp ty clauses') ~(pair $ map name strings))))
235 |                       :< `(Control.Arrow.first ~(foldr1 (\t,t' => `(Control.Arrow.(\|/) ~t ~t')) conts')))
236 |             strings rest
237 |     arrowImpl' ts strings (ILet _ _ MW (UN $ Basic var) ty exp rest) = do
238 |       (usedExp,usedLater,strings') <- do
239 |         let usedExp = map (usedInExpr exp . name) strings
240 |         let usedLater = map (\v => v.name /= var && usedInDiagram rest v.name) strings
241 |         let strings' = MkAVar var ty :: zipFilter strings usedLater
242 |         pure (usedExp,usedLater,strings')
243 |       arrowImpl'
244 |         (ts :< `(Control.Arrow.arrow {a = ~(pairTy strings), b = ~(pairTy strings')}
245 |           ~(stringLam (zip strings $ zipWith (\x,y => x || y) usedExp usedLater)
246 |             (ILet EmptyFC EmptyFC MW (UN $ Basic var) ty exp (pair $ map name strings')))))
247 |         strings' rest
248 |     arrowImpl' ts strings (ILet fc {}) = failAt fc "Let binding must be unrestricted"
249 |     arrowImpl' ts strings `((~(mor) -< ~(inp)) >>= ~(rest)) =
250 |       case strings of
251 |         [] => arrowImplLam (ts :< `(Control.Arrow.arrow {a=_,b=_} (\_ => ~inp)) :< mor) [] rest
252 |         _ :: _ => do
253 |           (usedInp,usedLater,strings') <- do
254 |             let usedInp = map (usedInExpr inp . name) strings
255 |             let usedLater = map (usedInDiagram rest . name) strings
256 |             let strings' = zipFilter strings usedLater
257 |             pure (usedInp,usedLater,strings')
258 |           case strings' of
259 |             [] =>
260 |               arrowImplLam
261 |                 (ts :< `(Control.Arrow.arrow {a = ~(pairTy strings),b=_}
262 |                           ~(stringLam (zip strings usedInp) inp))
263 |                     :< mor)
264 |                 strings' rest
265 |             _ :: _ =>
266 |               arrowImplLam
267 |                 (ts :< `(Control.Arrow.arrow {a = ~(pairTy strings), b = Builtin.Pair _ ~(pairTy strings')}
268 |                   ~(stringLam (zip strings $ zipWith (\x,y => x || y) usedInp usedLater)
269 |                     `(Builtin.MkPair ~inp ~(pair $ map name strings'))))
270 |                     :< `(Control.Arrow.first ~mor))
271 |                 strings' rest
272 |     arrowImpl' ts strings `((~(mor) -< ~(inp)) >> ~(rest)) =
273 |       case strings of
274 |         [] => arrowImpl' (ts :<
275 |                 `(Control.Arrow.arrow {a=_,b=_} (\_ => ~inp)) :< mor :< `(Control.Arrow.arrow {a=_,b=_} Builtin.snd)) [] rest
276 |         _ :: _ => do
277 |           (usedInp,usedLater,strings') <- do
278 |             let usedInp = map (usedInExpr inp . name) strings
279 |             let usedLater = map (usedInDiagram rest . name) strings
280 |             let strings' = zipFilter strings usedLater
281 |             pure (usedInp,usedLater,strings')
282 |           case strings' of
283 |             [] =>
284 |               arrowImpl'
285 |                 (ts :< `(Control.Arrow.arrow {a = ~(pairTy strings), b = _}
286 |                           ~(stringLam (zip strings usedInp) inp))
287 |                     :< mor :< `(Control.Arrow.arrow {a=_,b=_} (\_ => ())))
288 |                 strings' rest
289 |             _ :: _ =>
290 |               arrowImpl'
291 |                 (ts :< `(Control.Arrow.arrow {a = ~(pairTy strings), b = Builtin.Pair _ ~(pairTy strings')}
292 |                         ~(stringLam (zip strings $ zipWith (\x,y => x || y) usedInp usedLater)
293 |                           `(Builtin.MkPair ~inp ~(pair $ map name strings'))))
294 |                     :< `(Control.Arrow.first ~mor) :< `(Control.Arrow.arrow {a=_,b=_} Builtin.snd))
295 |                 strings' rest
296 |     arrowImpl' ts strings `(=< ~(out)) =
297 |       pure (ts :< `(Control.Arrow.arrow {a = ~(pairTy strings),b=_}
298 |         ~(stringLam (map (,True) strings) out)))
299 |
300 |     arrowImpl' ts strings (ILam fc {}) = failAt (getFC t) "Lambda not allowed here"
301 |     arrowImpl' ts strings t = failAt (getFC t) "Could not parse expression"
302 |
303 | ||| Enter arrow notation.
304 | |||
305 | ||| This elaboration script must be specifically invoked with
306 | ||| `%runElab`. The category is inferred from the return type.
307 | export
308 | arrowDo : {0 arr : Hom Type0} -> Arrow arr => TTImp -> Elab (arr (W0 a) (W0 b))
309 | arrowDo t = check !(arrowImpl t)
310 |