summaryrefslogtreecommitdiffstats
path: root/Map.sml
blob: a60c5cc59bde8ee7282496abd2de710b0db4be62 (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
(* 🅱️-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