summaryrefslogtreecommitdiffstats
path: root/Map.sml
blob: d8b3a377d62b2ef156d7366cdabc0c5cf36df453 (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
(* 🅱️-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 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 height (Tip : 'a map) : int = 0
    | height (Two (h, _, _, _, _)) = h
    | height (Three (h, _, _, _, _, _, _, _)) = h

  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 + 1, l, k, v, r)

  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 + 1, a, k1, v1, b, k2, v2, c)

  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 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"
      | 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 (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 => (Tip, false, Tip)
    | 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, Tip) => (l, (k, v))
    | SOME (l, k, v, r) =>
        let val (t', k') = splitLast r
        in (join l k v t', k')
        end
    | NONE => raise Fail "splitLast Tip"

  fun join2 (Tip : 'a map) (tr : 'a map) : 'a map = tr
    | join2 tl tr =
        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 (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 : '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 : '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 : (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 (Tip, k, v, _) => SOME (k, v)
    | SOME (l, _, _, _) => lookupMin l
end