jon.recoil.org

Source file jkind_axis.ml

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
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
(**************************************************************************)
(*                                                                        *)
(*                                 OCaml                                  *)
(*                                                                        *)
(*                  Liam Stevenson, Jane Street, New York                 *)
(*                                                                        *)
(*   Copyright 2024 Jane Street Group LLC                                 *)
(*                                                                        *)
(*   All rights reserved.  This file is distributed under the terms of    *)
(*   the GNU Lesser General Public License version 2.1, with the          *)
(*   special exception on linking described in the file LICENSE.          *)
(*                                                                        *)
(**************************************************************************)

module Fmt = Format_doc

module type Axis_ops = sig
  include Mode_intf.Lattice

  val to_string : t -> string

  val less_or_equal : t -> t -> Misc.Le_result.t

  val equal : t -> t -> bool
end

module Externality = struct
  type t =
    | External
    | External64
    | Internal

  include Mode.Lattices.Total (struct
    type nonrec t = t

    let min = External

    let max = Internal

    let ord = function External -> 0 | External64 -> 1 | Internal -> 2
  end)

  let less_or_equal s1 s2 : Misc.Le_result.t =
    if equal s1 s2 then Equal else if le s1 s2 then Less else Not_le

  let to_string = function
    | External -> "external_"
    | External64 -> "external64"
    | Internal -> "internal"

  let print ppf t = Fmt.fprintf ppf "%s" (to_string t)

  let upper_bound_if_is_always_gc_ignorable () =
    (* We check that we're compiling to (64-bit) native code before counting
        External64 types as gc_ignorable, because bytecode is intended to be
        platform independent. *)
    if !Clflags.native_code && Sys.word_size = 64 then External64 else External
end

module Nullability = struct
  type t =
    | Non_null
    | Maybe_null

  include Mode.Lattices.Total (struct
    type nonrec t = t

    let min = Non_null

    let max = Maybe_null

    let ord = function Non_null -> 0 | Maybe_null -> 1
  end)

  let less_or_equal s1 s2 : Misc.Le_result.t =
    if equal s1 s2 then Equal else if le s1 s2 then Less else Not_le

  let to_string = function Non_null -> "non_null" | Maybe_null -> "maybe_null"

  let print ppf t = Fmt.fprintf ppf "%s" (to_string t)
end

module Separability = struct
  type t =
    | Non_pointer
    | Non_pointer64
    | Non_float
    | Separable
    | Maybe_separable

  include Mode.Lattices.Total (struct
    type nonrec t = t

    let min = Non_pointer

    let max = Maybe_separable

    let ord = function
      | Non_pointer -> 0
      | Non_pointer64 -> 1
      | Non_float -> 2
      | Separable -> 3
      | Maybe_separable -> 4
  end)

  let less_or_equal s1 s2 : Misc.Le_result.t =
    if equal s1 s2 then Equal else if le s1 s2 then Less else Not_le

  let to_string = function
    | Non_pointer -> "non_pointer"
    | Non_pointer64 -> "non_pointer64"
    | Non_float -> "non_float"
    | Separable -> "separable"
    | Maybe_separable -> "maybe_separable"

  let print ppf t = Fmt.fprintf ppf "%s" (to_string t)

  let upper_bound_if_is_always_gc_ignorable () =
    (* We check that we're compiling to (64-bit) native code before counting
        Non_pointer64 types as gc_ignorable, because bytecode is intended to be
        platform independent. *)
    if !Clflags.native_code && Sys.word_size = 64
    then Non_pointer64
    else Non_pointer
end

module Axis = struct
  module Nonmodal = struct
    type 'a t = Externality : Externality.t t
  end

  type 'a t =
    | Modal : 'a Mode.Crossing.Axis.t -> 'a t
    | Nonmodal : 'a Nonmodal.t -> 'a t

  type packed = Pack : 'a t -> packed [@@unboxed]

  let all =
    [ Pack (Modal (Comonadic Areality));
      Pack (Modal (Monadic Uniqueness));
      Pack (Modal (Comonadic Linearity));
      Pack (Modal (Monadic Contention));
      Pack (Modal (Comonadic Portability));
      Pack (Modal (Comonadic Forkable));
      Pack (Modal (Comonadic Yielding));
      Pack (Modal (Comonadic Statefulness));
      Pack (Modal (Monadic Visibility));
      Pack (Modal (Monadic Staticity));
      (* CR-soon zqian: call [Mode.Crossing.Axis.all] for modal axes *)
      Pack (Nonmodal Externality) ]

  let name (type a) : a t -> string = function
    | Modal ax ->
      let (P ax) =
        P ax |> Mode.Crossing.Axis.to_modality |> Mode.Modality.Axis.to_value
      in
      Fmt.asprintf "%a" Mode.Value.Axis.print ax
    | Nonmodal Externality -> "externality"
end

module Per_axis = struct
  open Axis

  module Nonmodal = struct
    open Axis.Nonmodal

    let min : type a. a t -> a = function Externality -> Externality.min

    let max : type a. a t -> a = function Externality -> Externality.max

    let le : type a. a t -> a -> a -> bool =
     fun ax a b -> match ax with Externality -> Externality.le a b

    let equal : type a. a t -> a -> a -> bool =
     fun ax a b -> match ax with Externality -> Externality.equal a b

    let meet : type a. a t -> a -> a -> a =
     fun ax a b -> match ax with Externality -> Externality.meet a b

    let join : type a. a t -> a -> a -> a =
     fun ax a b -> match ax with Externality -> Externality.join a b

    let print : type a. a t -> Fmt.formatter -> a -> unit = function
      | Externality -> Externality.print

    let compare_obj : type a b. a t -> b t -> int =
     fun a b -> match a, b with Externality, Externality -> 0

    let equal_obj : type a b. a t -> b t -> (a, b) Misc.is_eq =
     fun a b -> match a, b with Externality, Externality -> Misc.Is_eq
  end

  let min : type a. a t -> a = function[@inline available]
    | Modal ax -> (Mode.Crossing.Per_axis.min [@inlined hint]) ax
    | Nonmodal ax -> (Nonmodal.min [@inlined hint]) ax

  let max : type a. a t -> a = function[@inline available]
    | Modal ax -> (Mode.Crossing.Per_axis.max [@inlined hint]) ax
    | Nonmodal ax -> (Nonmodal.max [@inlined hint]) ax

  let le : type a. a t -> a -> a -> bool =
   fun[@inline available] ax a b ->
    match ax with
    | Modal ax -> (Mode.Crossing.Per_axis.le [@inlined hint]) ax a b
    | Nonmodal ax -> (Nonmodal.le [@inlined hint]) ax a b

  let equal : type a. a t -> a -> a -> bool =
   fun[@inline available] ax a b ->
    match ax with
    | Modal ax -> (Mode.Crossing.Per_axis.equal [@inlined hint]) ax a b
    | Nonmodal ax -> (Nonmodal.equal [@inlined hint]) ax a b

  let meet : type a. a t -> a -> a -> a =
   fun[@inline available] ax a b ->
    match ax with
    | Modal ax -> (Mode.Crossing.Per_axis.meet [@inlined hint]) ax a b
    | Nonmodal ax -> (Nonmodal.meet [@inlined hint]) ax a b

  let join : type a. a t -> a -> a -> a =
   fun[@inline available] ax a b ->
    match ax with
    | Modal ax -> (Mode.Crossing.Per_axis.join [@inlined hint]) ax a b
    | Nonmodal ax -> (Nonmodal.join [@inlined hint]) ax a b

  let print : type a. a t -> Fmt.formatter -> a -> unit = function
    | Modal ax -> Mode.Crossing.Per_axis.print ax
    | Nonmodal ax -> Nonmodal.print ax

  let compare_obj : type a b. a t -> b t -> int =
   fun a b ->
    match a, b with
    | Modal ax0, Modal ax1 -> Mode.Crossing.Per_axis.compare_obj ax0 ax1
    | Modal _, _ -> -1
    | _, Modal _ -> 1
    | Nonmodal ax0, Nonmodal ax1 -> Nonmodal.compare_obj ax0 ax1

  let equal_obj : type a b. a t -> b t -> (a, b) Misc.is_eq =
   fun a b ->
    match a, b with
    | Modal ax0, Modal ax1 -> Mode.Crossing.Per_axis.equal_obj ax0 ax1
    | Modal _, _ -> Misc.Is_not_eq
    | _, Modal _ -> Misc.Is_not_eq
    | Nonmodal ax0, Nonmodal ax1 -> Nonmodal.equal_obj ax0 ax1

  let print_obj : type a. Fmt.formatter -> a t -> unit =
   fun ppf ax -> Fmt.pp_print_string ppf (name ax)
end

module Axis_set = struct
  (* This could be [bool Axis_collection.t], but instead we represent it as a bitfield for
     performance (this matters, since these are hammered on quite a bit during with-bound
     normalization) *)

  type t = int

  let[@inline] axis_index (type a) : a Axis.t -> _ = function
    | Modal (Comonadic Areality) -> 0
    | Modal (Monadic Uniqueness) -> 1
    | Modal (Comonadic Linearity) -> 2
    | Modal (Monadic Contention) -> 3
    | Modal (Comonadic Portability) -> 4
    | Modal (Comonadic Forkable) -> 5
    | Modal (Comonadic Yielding) -> 6
    | Modal (Comonadic Statefulness) -> 7
    | Modal (Monadic Visibility) -> 8
    | Modal (Monadic Staticity) -> 9
    (* CR-soon zqian: call [Mode.Crossing.Axis.index] for modal axes *)
    | Nonmodal Externality -> 10

  let[@inline] axis_mask ax = 1 lsl axis_index ax

  let[@inline] set ~axis ~to_ t =
    match to_ with
    | true -> t lor axis_mask axis
    | false -> t land lnot (axis_mask axis)

  let empty = 0

  let[@inline] add t axis = set ~axis ~to_:true t

  let[@inline] create ~f =
    (* PERF: this is manually unrolled because flambda2 doesn't unroll for us, and this
       function is quite hot *)
    let[@inline] set_axis axis t =
      if f ~axis:(Axis.Pack axis) then t lor axis_mask axis else t
    in
    0
    |> set_axis (Modal (Comonadic Areality))
    |> set_axis (Modal (Monadic Uniqueness))
    |> set_axis (Modal (Comonadic Linearity))
    |> set_axis (Modal (Monadic Contention))
    |> set_axis (Modal (Comonadic Portability))
    |> set_axis (Modal (Comonadic Forkable))
    |> set_axis (Modal (Comonadic Yielding))
    |> set_axis (Modal (Comonadic Statefulness))
    |> set_axis (Modal (Monadic Visibility))
    |> set_axis (Modal (Monadic Staticity))
    |> set_axis (Nonmodal Externality)

  let all = create ~f:(fun ~axis:_ -> true)

  let equal = Int.equal

  let all_modal_axes =
    create ~f:(fun ~axis ->
        match axis with Pack (Modal _) -> true | Pack (Nonmodal _) -> false)

  let[@inline] singleton axis = add empty axis

  let[@inline] remove t axis = set ~axis ~to_:false t

  let[@inline] mem t axis = not (Int.equal (t land axis_mask axis) 0)

  let[@inline] union t1 t2 = t1 lor t2

  let[@inline] intersection t1 t2 = t1 land t2

  let[@inline] diff t1 t2 = t1 land lnot t2

  let[@inline] is_subset t1 t2 = Int.equal (t1 land t2) t1

  let[@inline] is_empty t = Int.equal t 0

  let[@inline] complement t = diff all t

  let all_nonmodal_axes = complement all_modal_axes

  let[@inline] to_seq t =
    Axis.all |> List.to_seq |> Seq.filter (fun (Axis.Pack axis) -> mem t axis)

  let[@inline] to_list t = List.of_seq (to_seq t)

  let print ppf t =
    Format.fprintf ppf "@[{%t}@]" (fun ppf ->
        Format.pp_print_seq
          ~pp_sep:(fun ppf () -> Format.fprintf ppf ";@ ")
          (fun ppf (Axis.Pack axis) -> Format.fprintf ppf "%s" (Axis.name axis))
          ppf (to_seq t))
end