diff options
Diffstat (limited to 'map.sml')
| -rw-r--r-- | map.sml | 254 |
1 files changed, 123 insertions, 131 deletions
@@ -1,6 +1,4 @@ -(* map is a size-balanced binary tree implementing an immutable key-value store. - - * Stolen with love from Haskell's Data.Map. *) +(* 🅱️-tree *) signature MAP = sig type key @@ -8,7 +6,6 @@ sig 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 @@ -25,170 +22,165 @@ functor Map (type k struct type key = k - datatype 'a map = Tip - | Bin of int * key * 'a * 'a map * 'a map + datatype 'a map = + Tip + | Two of int * 'a map * key * 'a * 'a map + | Three of int * 'a map * key * 'a * 'a map * key * 'a * '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 height (Tip : 'a map) : int = 0 + | height (Two (h, _, _, _, _)) = h + | height (Three (h, _, _, _, _, _, _, _)) = h - 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 two (l : 'a map) (k : key) (v : 'a) (r : 'a map) : 'a map = + if height l <> height r then raise Fail "two: height mismatch" else + Two (height l, l, k, v, r) - 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 three (a : 'a map) (k1 : key) (v1 : 'a) (b : 'a map) (k2 : key) (v2 : 'a) (c : 'a map) : 'a map = + if not (height a = height b andalso height b = height c) then raise Fail "three: height mismatch" else + Three (height a, a, k1, v1, b, k2, v2, c) - 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 view (Tip : 'a map) : ('a map * key * 'a * 'a map) option = NONE + | view (Two (_, l, k, v, r)) = SOME (l, k, v, r) + | view (Three (h, a, k1, v1, b, k2, v2, c)) = SOME (a, k1, v1, two b k2 v2 c) - 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 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 - 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) + datatype 'a insertResult = + One of 'a map + | Split of 'a map * key * 'a * 'a map - 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' (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" + | Two (_, rl, rk, rv, rr) => + (case join' left k v rl of + One newNode => One (two newNode rk rv rr) + | Split (left, k, v, right) => One (three left k v right rk rv rr)) + | Three (_, ra, rk1, rv1, rb, rk2, rv2, rc) => + case join' left k v ra of + One newNode => One (three newNode rk1 rv1 rb rk2 rv2 rc) + | Split (left, k, v, right) => Split (two left k v right, rk1, rv1, two rb rk2 rv2 rc) + else + case left of + Tip => raise Fail "unreachable" + | Two (_, ll, lk, lv, lr) => + (case join' lr k v right of + One newNode => One (two ll lk lv newNode) + | Split (left, k, v, right) => One (three ll lk lv left k v right)) + | Three (_, la, lk1, lv1, lb, lk2, lv2, lc) => + case join' lc k v right of + One newNode => One (three la lk1 lv1 lb lk2 lv2 newNode) + | Split (left, k, v, right) => Split (two la lk1 lv1 lb, lk2, lv2, two left k v right) - 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 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 (t : 'a map) (k : key) = - case t of - Tip => (Tip, false, Tip) - | Bin (_, km, m, l, r) => - case cmp (k, km) of + fun split (m : 'a map) (k : key) : 'a map * bool * 'a map = + case view m of + NONE => (Tip, false, Tip) + | SOME (l, k', v, r) => + case cmp (k, k') of EQUAL => (l, true, r) | LESS => - let val (ll, b, lr) = split l k - in (ll, b, join km m lr r) + let val (ll, found, lr) = split l k + in (ll, found, join lr k' v r) end | GREATER => - let val (rl, b, rr) = split r k - in (join km m l rl, b, rr) + let val (rl, found, rr) = split r k + in (join l k' v rl, found, rr) end - fun splitLast (Bin (_, kx, x, l, Tip)) = (l, (kx, x)) - | splitLast (Bin (_, kx, x, l, r)) = + fun splitLast (m : 'a map) : 'a map * (key * 'a) = + case view m of + SOME (l, k, v, Tip) => (l, (k, v)) + | SOME (l, k, v, r) => let val (t', k') = splitLast r - in (join kx x l t', k') + in (join l k v t', k') end - | splitLast Tip = raise Fail "splitLast Tip" + | NONE => raise Fail "splitLast Tip" - fun join2 Tip tr = tr + fun join2 (Tip : 'a map) (tr : 'a map) : 'a map = tr | join2 tl tr = let val (tl', (kx, x)) = splitLast tl - in join kx x tl' tr + in join tl' kx x tr end - fun delete k m = + 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 = - 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 union (Tip : 'a map) (t2 : 'a map) : 'a map = t2 + | union t1 t2 = + 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 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 intersect (Tip : 'a map) (_ : 'a map) : 'a map = Tip + | intersect t1 t2 = + case view t2 of + NONE => Tip + | 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 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 difference (Tip : 'a map) (_ : 'b map) : 'a map = Tip + | difference t1 t2 = + 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 Tip k v Tip 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 fromList (l : (key * 'a) list) : 'a map = 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 : '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 = toList' m [] + fun toList (m : 'a map) : (key * 'a) list = toList' m [] - fun lookupMin Tip = NONE - | lookupMin (Bin (_, k, v, Tip, _)) = SOME (k, v) - | lookupMin (Bin (_, _, _, l, _)) = lookupMin l + fun lookupMin (m : 'a map) : (key * 'a) option = + case view m of + NONE => NONE + | SOME (Tip, k, v, _) => SOME (k, v) + | SOME (l, _, _, _) => lookupMin l end |
