summaryrefslogtreecommitdiffstats
path: root/map.sml
blob: a5b316c942f7dc1026d4b101c4ffcce3e41b1f08 (plain) (blame)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
(* 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