(* 🅱️-tree *) signature MAP = sig type key type 'a map val empty : 'a map val null : 'a map -> bool val insert : key -> 'a -> 'a map -> 'a map val lookup : key -> 'a map -> 'a option val delete : key -> 'a map -> 'a map val union : 'a map -> 'a map -> 'a map val intersect : 'a map -> 'a map -> 'a map val difference : 'a map -> 'b map -> 'a map val fromList : (key * 'a) list -> 'a map val toList : 'a map -> (key * 'a) list val lookupMin : 'a map -> (key * 'a) option end functor Map (type k val cmp : k * k -> order) :> MAP where type key = k = struct type key = k datatype 'a node = Tip | Two of 'a node * key * 'a * 'a node | Three of 'a node * key * 'a * 'a node * key * 'a * 'a node type 'a map = int * 'a node val empty : 'a map = (0, Tip) fun height ((h, _) : 'a map) : int = h fun null (m : 'a map) : bool = height m = 0 fun two ((lh, l) : 'a map) (k : key) (v : 'a) ((rh, r) : 'a map) : 'a map = if lh <> rh then raise Fail "two: height mismatch" else (lh + 1, Two (l, k, v, r)) fun three ((ah, a) : 'a map) (k1 : key) (v1 : 'a) ((bh, b) : 'a map) (k2 : key) (v2 : 'a) ((ch, c) : 'a map) : 'a map = if not (ah = bh andalso bh = ch) then raise Fail "three: height mismatch" else (ah + 1, Three (a, k1, v1, b, k2, v2, c)) fun view ((_, Tip) : 'a map) : ('a map * key * 'a * 'a map) option = NONE | view (h, Two (l, k, v, r)) = SOME ((h - 1, l), k, v, (h - 1, r)) | view (h, Three (a, k1, v1, b, k2, v2, c)) = SOME ((h - 1, a), k1, v1, two (h - 1, b) k2 v2 (h - 1, c)) fun lookup (k : key) (m : 'a map) : 'a option = case view m of NONE => NONE | SOME (l, k', v, r) => case cmp (k, k') of EQUAL => SOME v | LESS => lookup k l | GREATER => lookup k r datatype 'a insertResult = One of 'a map | Split of 'a map * key * 'a * 'a map fun join' (left : 'a map) (k : key) (v : 'a) (right : 'a map) : 'a insertResult = if height left = height right then Split (left, k, v, right) else if height left < height right then case right of (_, Tip) => raise Fail "unreachable" | (rh, Two (rl, rk, rv, rr)) => (case join' left k v (rh - 1, rl) of One newNode => One (two newNode rk rv (rh - 1, rr)) | Split (left, k, v, right) => One (three left k v right rk rv (rh - 1, rr))) | (rh, Three (ra, rk1, rv1, rb, rk2, rv2, rc)) => case join' left k v (rh - 1, ra) of One newNode => One (three newNode rk1 rv1 (rh - 1, rb) rk2 rv2 (rh - 1, rc)) | Split (left, k, v, right) => Split (two left k v right, rk1, rv1, two (rh - 1, rb) rk2 rv2 (rh - 1, rc)) else case left of (_, Tip) => raise Fail "unreachable" | (lh, Two (ll, lk, lv, lr)) => (case join' (lh - 1, lr) k v right of One newNode => One (two (lh - 1, ll) lk lv newNode) | Split (left, k, v, right) => One (three (lh - 1, ll) lk lv left k v right)) | (lh, Three (la, lk1, lv1, lb, lk2, lv2, lc)) => case join' (lh - 1, lc) k v right of One newNode => One (three (lh - 1, la) lk1 lv1 (lh - 1, lb) lk2 lv2 newNode) | Split (left, k, v, right) => Split (two (lh - 1, la) lk1 lv1 (lh - 1, lb), lk2, lv2, two left k v right) fun join (left : 'a map) (k : key) (v : 'a) (right : 'a map) : 'a map = case join' left k v right of One node => node | Split (left, k, v, right) => two left k v right fun split (m : 'a map) (k : key) : 'a map * bool * 'a map = case view m of NONE => (empty, false, empty) | SOME (l, k', v, r) => case cmp (k, k') of EQUAL => (l, true, r) | LESS => let val (ll, found, lr) = split l k in (ll, found, join lr k' v r) end | GREATER => let val (rl, found, rr) = split r k in (join l k' v rl, found, rr) end fun splitLast (m : 'a map) : 'a map * (key * 'a) = case view m of SOME (l, k, v, r) => if null r then (l, (k, v)) else let val (t', k') = splitLast r in (join l k v t', k') end | NONE => raise Fail "splitLast Tip" fun join2 (tl : 'a map) (tr : 'a map) : 'a map = if null tl then tr else let val (tl', (kx, x)) = splitLast tl in join tl' kx x tr end fun delete (k : key) (m : 'a map) : 'a map = let val (tl, _, tr) = split m k in join2 tl tr end fun union (t1 : 'a map) (t2 : 'a map) : 'a map = if null t1 then t2 else case view t2 of NONE => t1 | SOME (l2, k2, v2, r2) => let val (l1, _, r1) = split t1 k2 val tl = union l1 l2 val tr = union r1 r2 in join tl k2 v2 tr end fun intersect (t1 : 'a map) (t2 : 'a map) : 'a map = if null t1 then empty else case view t2 of NONE => empty | SOME (l2, k2, v2, r2) => let val (l1, b, r1) = split t1 k2 val tl = intersect l1 l2 val tr = intersect r1 r2 in if b then join tl k2 v2 tr else join2 tl tr end fun difference (t1 : 'a map) (t2 : 'b map) : 'a map = if null t1 then empty else case view t2 of NONE => t1 | SOME (l2, k2, _, r2) => let val (l1, _, r1) = split t1 k2 val tl = difference l1 l2 val tr = difference r1 r2 in join2 tl tr end fun singleton (k : key) (v : 'a) : 'a map = two empty k v empty fun insert (k : key) (v : 'a) (m : 'a map) : 'a map = union m (singleton k v) fun fromList (l : (key * 'a) list) : 'a map = foldl (fn ((kx, x), acc) => insert kx x acc) empty l fun toList' (m : 'a map) (acc : (key * 'a) list) : (key * 'a) list = case view m of NONE => acc | SOME (l, k, v, r) => toList' l ((k, v) :: toList' r acc) fun toList (m : 'a map) : (key * 'a) list = toList' m [] fun lookupMin (m : 'a map) : (key * 'a) option = case view m of NONE => NONE | SOME (l, k, v, _) => if null l then SOME (k, v) else lookupMin l end