(* map is a size-balanced binary tree implementing an immutable key-value store. * Stolen with love from Haskell's Data.Map. *) signature MAP = sig type key type 'a map val empty : 'a map val null : 'a map -> bool val size : 'a map -> int 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 map = Tip | Bin of int * key * 'a * 'a map * 'a map val empty : 'a map = Tip fun null Tip = true | null _ = false fun size Tip = 0 | size (Bin (s, _, _, _, _)) = s fun lookup (k : key) (m : 'a map) : 'a option = let fun go Tip = NONE | go (Bin (_, kx, x, l, r)) = case cmp (k, kx) of LESS => go l | GREATER => go r | EQUAL => SOME x in go m end fun singleton (k : key) (v : 'a) : 'a map = Bin (1, k, v, Tip, Tip) fun size Tip = 0 | size (Bin (sz, _, _, _, _)) = sz (* The bin constructor maintains the size of the tree. *) fun bin k x l r = Bin (size l + size r + 1, k, x, l, r) (* https://hackage.haskell.org/package/containers-0.4.0.0/docs/src/Data-Map.html#delta *) val delta = 4 val ratio = 2 fun singleL k1 x1 t1 (Bin (_, k2, x2, t2, t3)) = bin k2 x2 (bin k1 x1 t1 t2) t3 | singleL _ _ _ Tip = raise Fail "singleL Tip" fun singleR k1 x1 (Bin (_, k2, x2, t1, t2)) t3 = bin k2 x2 t1 (bin k1 x1 t2 t3) | singleR _ _ Tip _ = raise Fail "singleR Tip" fun doubleL k1 x1 t1 (Bin (_, k2, x2, Bin (_, k3, x3, t2, t3), t4)) = bin k3 x3 (bin k1 x1 t1 t2) (bin k2 x2 t3 t4) | doubleL _ _ _ _ = raise Fail "doubleL" fun doubleR k1 x1 (Bin (_, k2, x2, t1, Bin (_, k3, x3, t2, t3))) t4 = bin k3 x3 (bin k2 x2 t1 t2) (bin k1 x1 t3 t4) | doubleR _ _ _ _ = raise Fail "doubleR" fun rotateL k x l (r as Bin (_, _, _, ly, ry)) = if size ly < ratio*size ry then singleL k x l r else doubleL k x l r | rotateL _ _ _ Tip = raise Fail "rotateL Tip" fun rotateR k x (l as Bin (_, _, _, ly, ry)) r = if size ry < ratio*size ly then singleR k x l r else doubleR k x l r | rotateR _ _ Tip _ = raise Fail "rotateR Tip" fun balance k x l r = let val sizeL = size l val sizeR = size r in if sizeL + sizeR <= 1 then bin k x l r else if sizeR >= delta*sizeL then rotateL k x l r else if sizeL >= delta*sizeR then rotateR k x l r else bin k x l r end fun insertMax kx x t = case t of Tip => singleton kx x | Bin (_, ky, y, l, r) => balance ky y l (insertMax kx x r) fun insertMin kx x t = case t of Tip => singleton kx x | Bin (_, ky, y, l, r) => balance ky y (insertMin kx x l) r fun join kx x Tip r = insertMin kx x r | join kx x l Tip = insertMax kx x l | join kx x (l as Bin (sizeL, ky, y, ly, ry)) (r as Bin (sizeR, kz, z, lz, rz)) = if delta*sizeL <= sizeR then balance kz z (join kx x l lz) rz else if delta*sizeR <= sizeL then balance ky y ly (join kx x ry r) else bin kx x l r fun split (t : 'a map) (k : key) = case t of Tip => (Tip, false, Tip) | Bin (_, km, m, l, r) => case cmp (k, km) of EQUAL => (l, true, r) | LESS => let val (ll, b, lr) = split l k in (ll, b, join km m lr r) end | GREATER => let val (rl, b, rr) = split r k in (join km m l rl, b, rr) end fun splitLast (Bin (_, kx, x, l, Tip)) = (l, (kx, x)) | splitLast (Bin (_, kx, x, l, r)) = let val (t', k') = splitLast r in (join kx x l t', k') end | splitLast Tip = raise Fail "splitLast Tip" fun join2 Tip tr = tr | join2 tl tr = let val (tl', (kx, x)) = splitLast tl in join kx x tl' tr end fun delete k m = let val (tl, _, tr) = split m k in join2 tl tr end fun union (t1 : 'a map) (t2 : 'a map) : 'a map = case (t1, t2) of (Tip, _) => t2 | (_, Tip) => t1 | (_, Bin (_, k2, v2, l2, r2)) => let val (l1, _, r1) = split t1 k2 val tl = union l1 l2 val tr = union r1 r2 in join k2 v2 tl tr end fun intersect Tip _ = Tip | intersect _ Tip = Tip | intersect t1 (Bin (_, k2, v2, l2, r2)) = let val (l1, b, r1) = split t1 k2 val tl = intersect l1 l2 val tr = intersect r1 r2 in if b then join k2 v2 tl tr else join2 tl tr end fun difference Tip _ = Tip | difference t1 Tip = t1 | difference t1 (Bin (_, k2, v2, l2, r2)) = let val (l1, _, r1) = split t1 k2 val tl = difference l1 l2 val tr = difference r1 r2 in join2 tl tr end fun insert (k : key) (v : 'a) (m : 'a map) : 'a map = union m (singleton k v) fun fromList l = foldl (fn ((kx, x), acc) => insert kx x acc) empty l fun toList' Tip acc = acc | toList' (Bin (_, k, v, l, r)) acc = toList' l ((k, v) :: toList' r acc) fun toList m = toList' m [] fun lookupMin Tip = NONE | lookupMin (Bin (_, k, v, Tip, _)) = SOME (k, v) | lookupMin (Bin (_, _, _, l, _)) = lookupMin l end