From cd9d7ca6b63260e4426b3df8291891293b403609 Mon Sep 17 00:00:00 2001 From: Simon Cruanes Date: Wed, 30 Sep 2026 14:18:10 -0400 Subject: [PATCH 1/5] add main-domain check to ExtThread --- extThread.ml | 2 ++ extThread.mli | 4 ++++ extThreadBase.ml | 7 +++++++ log.ml | 16 ++++++---------- 4 files changed, 19 insertions(+), 10 deletions(-) create mode 100644 extThreadBase.ml diff --git a/extThread.ml b/extThread.ml index 68e1b625..9a1ae1eb 100644 --- a/extThread.ml +++ b/extThread.ml @@ -1,3 +1,5 @@ +include ExtThreadBase + let log = Log.self type 'a t = [ `Exn of exn | `None | `Ok of 'a ] ref * Thread.t diff --git a/extThread.mli b/extThread.mli index b366ae70..52f1f5d4 100644 --- a/extThread.mli +++ b/extThread.mli @@ -1,5 +1,9 @@ (** Thread utilities *) +val check_main_domain : string -> unit +(** [check_main_domain name] raises [Failure] if not called from the main domain. + Use it to guard process-wide setup and configuration. [name] identifies the caller in the error. *) + val locked : Mutex.t -> (unit -> 'a) -> 'a type 'a t diff --git a/extThreadBase.ml b/extThreadBase.ml new file mode 100644 index 00000000..9fa7c79e --- /dev/null +++ b/extThreadBase.ml @@ -0,0 +1,7 @@ +(** Thread and domain utilities with no devkit dependencies, see {!ExtThread} *) + +(** [check_main_domain name] raises [Failure] if not called from the main domain. + Use it to guard process-wide setup and configuration. [name] identifies the caller in the error. *) +let check_main_domain name = + if (Domain.self () :> int) <> 0 then + failwith (name ^ ": must be called from the main domain") diff --git a/log.ml b/log.ml index 0c136853..e0da2b9d 100644 --- a/log.ml +++ b/log.ml @@ -55,17 +55,13 @@ open Prelude (** Global logger state *) module State = struct - let check_main_domain name = - if not (Domain.is_main_domain ()) then - Exn.fail "Log.%s: must be called from the main domain" name - let all = Hashtbl.create 10 let default_level = ref (`Info : Logger.level) let utc_timezone = ref false let facility name = - check_main_domain "facility"; + ExtThreadBase.check_main_domain "Log.facility"; try Hashtbl.find all name with @@ -75,7 +71,7 @@ module State = struct x let set_filter ?name level = - check_main_domain "set_filter"; + ExtThreadBase.check_main_domain "Log.set_filter"; match name with | None -> default_level := level; Hashtbl.iter (fun _ x -> Logger.set_filter x level) all | Some name when Stre.ends_with name "*" -> @@ -84,7 +80,7 @@ module State = struct | Some name -> Logger.set_filter (facility name) level let set_loglevels s = - check_main_domain "set_loglevels"; + ExtThreadBase.check_main_domain "Log.set_loglevels"; Stre.nsplitc s ',' |> List.iter begin fun spec -> match Stre.nsplitc spec '=' with | name :: l :: [] -> set_filter ~name (Logger.level l) @@ -136,8 +132,8 @@ module State = struct end let get_cur_format () = Atomic.get cur_format let is_structured_format () = match get_cur_format () with `Plain, _ -> false | `Logfmt, _ -> true - let set_plaintext () = check_main_domain "set_plaintext"; set_cur_format (`Plain, format_simple_full) - let set_logfmt () = check_main_domain "set_logfmt"; set_cur_format (`Logfmt, format_logfmt) + let set_plaintext () = ExtThreadBase.check_main_domain "Log.set_plaintext"; set_cur_format (`Plain, format_simple_full) + let set_logfmt () = ExtThreadBase.check_main_domain "Log.set_logfmt"; set_cur_format (`Logfmt, format_logfmt) let format level facil ts pairs msg = (snd (Atomic.get cur_format)) level facil ts pairs msg @@ -184,7 +180,7 @@ end let facility = State.facility let set_filter = State.set_filter let set_loglevels = State.set_loglevels -let set_utc () = State.check_main_domain "set_utc"; State.utc_timezone := true +let set_utc () = ExtThreadBase.check_main_domain "Log.set_utc"; State.utc_timezone := true (** Update facilities configuration from the environment. From ab1fcf261c6a1c17abeabb5128b61c1666effb67 Mon Sep 17 00:00:00 2001 From: Simon Cruanes Date: Wed, 30 Sep 2026 22:36:24 -0400 Subject: [PATCH 2/5] fast lockfree map in ExtThread to be used in Cache. --- extThread.ml | 117 ++++++++++++++++++++++++++++++++++++++++++++++++++ extThread.mli | 56 ++++++++++++++++++++++++ 2 files changed, 173 insertions(+) diff --git a/extThread.ml b/extThread.ml index 9a1ae1eb..0bfda128 100644 --- a/extThread.ml +++ b/extThread.ml @@ -175,3 +175,120 @@ module Pool = struct end end + +(* Writes copy the path from the shard root to the leaf, then CAS the root. *) +module ShardedHashTrie = struct + type 'k hashable = { equal : 'k -> 'k -> bool; hash : 'k -> int } + + (* [compare], not [(=)], as in Hashtbl: [nan] must be equal to itself *) + let poly_hashable = { equal = (fun a b -> compare a b = 0); hash = Hashtbl.hash } + + let level_bits = 4 + let level_width = 1 lsl level_bits + let level_mask = level_width - 1 + + (* assoc list (all keys have same hash) *) + type ('k, 'v) bindings = + | Nil + | Cons of { key : 'k; value : 'v; rest : ('k, 'v) bindings } + + type ('k, 'v) tree = + | Empty + | Leaf of { h : int; key : 'k; value : 'v; next : ('k, 'v) bindings } + (** [h]: hash bits remaining at this depth; [next]: other keys with the same hash *) + | Node of ('k, 'v) tree array (** [level_width] children *) + + type ('k, 'v) t = { + hashable : 'k hashable; + shard_bits : int; + shard_mask : int; + shards : ('k, 'v) tree Atomic.t array; + } + + let create ?(hashable = poly_hashable) ?(shard_bits = 6) () = + if shard_bits < 0 || shard_bits > 16 then invalid_arg "ShardedHashTrie.create: shard_bits must be in [0,16]"; + { hashable; shard_bits; shard_mask = 1 lsl shard_bits - 1; + shards = Array.init (1 lsl shard_bits) (fun _ -> Atomic.make Empty) } + + let rec assoc equal k = function + | Nil -> raise_notrace Not_found + | Cons { key; value; rest } -> if equal k key then value else assoc equal k rest + + let rec find_tree equal h k = function + | Empty -> raise_notrace Not_found + | Leaf { h = h2; key; value; next } -> + if h <> h2 then raise_notrace Not_found + else if equal k key then value + else assoc equal k next + | Node a -> find_tree equal (h lsr level_bits) k a.(h land level_mask) + + (* shard and remaining hash computed inline: a helper returning both would allocate a tuple *) + let find_notrace t k = + let h = t.hashable.hash k in + find_tree t.hashable.equal (h lsr t.shard_bits) k (Atomic.get t.shards.(h land t.shard_mask)) + + (* re-raise so that callers get a backtrace *) + let find t k = try find_notrace t k with Not_found -> raise Not_found + + let find_opt t k = match find_notrace t k with v -> Some v | exception Not_found -> None + let mem t k = match find_notrace t k with _ -> true | exception Not_found -> false + + (* Pure: returns a new tree. [k] must not be bound in [tree]. *) + let rec insert h k v tree = + match tree with + | Empty -> Leaf { h; key = k; value = v; next = Nil } + | Leaf { h = h2; key; value; next } when h = h2 -> + Leaf { h; key = k; value = v; next = Cons { key; value; rest = next } } + | Leaf { h = h2; key; value; next } -> + (* different hashes: replace the leaf with a node holding it one level + down, and insert there. Terminates since [h] and [h2] differ in some bit. *) + let a = Array.make level_width Empty in + a.(h2 land level_mask) <- Leaf { h = h2 lsr level_bits; key; value; next }; + let i = h land level_mask in + a.(i) <- insert (h lsr level_bits) k v a.(i); + Node a + | Node a -> + let i = h land level_mask in + let a = Array.copy a in + a.(i) <- insert (h lsr level_bits) k v a.(i); + Node a + + let get_or_create t k ~f = + let h = t.hashable.hash k in + let shard = t.shards.(h land t.shard_mask) in + let h = h lsr t.shard_bits in + let equal = t.hashable.equal in + let tree = Atomic.get shard in + match find_tree equal h k tree with + | v -> v + | exception Not_found -> + let v = f k in + (* invariant: [k] is not bound in [tree] *) + let rec add tree = + if Atomic.compare_and_set shard tree (insert h k v tree) then v + else + let tree = Atomic.get shard in + match find_tree equal h k tree with + | v' -> v' (* another domain bound [k] first: drop [v] *) + | exception Not_found -> add tree + in + add tree + + let rec fold_bindings f acc = function + | Nil -> acc + | Cons { key; value; rest } -> fold_bindings f (f key value acc) rest + + let rec fold_tree f acc = function + | Empty -> acc + | Leaf { key; value; next; _ } -> fold_bindings f (f key value acc) next + | Node a -> Array.fold_left (fold_tree f) acc a + + let fold t f acc = + let snapshot = Array.map Atomic.get t.shards in + Array.fold_left (fold_tree f) acc snapshot + + let iter t f = fold t (fun k v () -> f k v) () + let to_list t = fold t (fun k v acc -> (k, v) :: acc) [] + let length t = fold t (fun _ _ n -> n + 1) 0 + let clear t = Array.iter (fun shard -> Atomic.set shard Empty) t.shards +end diff --git a/extThread.mli b/extThread.mli index 52f1f5d4..df34b72d 100644 --- a/extThread.mli +++ b/extThread.mli @@ -4,6 +4,62 @@ val check_main_domain : string -> unit (** [check_main_domain name] raises [Failure] if not called from the main domain. Use it to guard process-wide setup and configuration. [name] identifies the caller in the error. *) +(** Domain-safe hash table, fast for reads. + + All operations may be called concurrently from any domain. + Lookups are wait-free and do not allocate. Writes are lock-free: they copy + a short path of the underlying trie and retry on contention, so inserts are + about twice as slow as with [Hashtbl]. Meant for read-mostly tables, eg. a + set of keys that is looked up far more often than it grows. + + There is no resizing: the table is split into a fixed number of shards + (see [shard_bits] in {!create}), each holding a hash trie of branching + factor 16 that deepens as it grows. *) +module ShardedHashTrie : sig + type 'k hashable = { equal : 'k -> 'k -> bool; hash : 'k -> int } + (** [equal a b] implies [hash a = hash b]. Keys with equal hashes are kept in a + list, so a hash with few distinct values makes the table slow. *) + + val poly_hashable : 'k hashable + (** Same as [Hashtbl]: [compare a b = 0] and [Hashtbl.hash]. For string keys, prefer + [{ equal = String.equal; hash = Hashtbl.hash }]. *) + + type ('k, 'v) t + + val create : ?hashable:'k hashable -> ?shard_bits:int -> unit -> ('k, 'v) t + (** @param hashable defaults to {!poly_hashable} + @param shard_bits log2 of the number of shards, in [\[0,16\]]. Default 6 + (64 shards, about 1.5KB), fine for up to a few thousand keys; use 8 or more for larger tables. + @raise Invalid_argument if [shard_bits] is out of range *) + + val find : ('k, 'v) t -> 'k -> 'v + (** @raise Not_found if the key is not bound *) + + val find_opt : ('k, 'v) t -> 'k -> 'v option + val mem : ('k, 'v) t -> 'k -> bool + + val get_or_create : ('k, 'v) t -> 'k -> f:('k -> 'v) -> 'v + (** Return the value bound to the key, binding it to [f key] first if absent. + Concurrent callers for the same absent key may each call [f], but only one + result is stored and all of them get that one. *) + + (** {2 Iteration} + + Iteration works on a snapshot of the shards taken at the start, so writes + made during iteration are not observed. The snapshot is not atomic across + shards. Order is unspecified. *) + + val iter : ('k, 'v) t -> ('k -> 'v -> unit) -> unit + val fold : ('k, 'v) t -> ('k -> 'v -> 'acc -> 'acc) -> 'acc -> 'acc + val to_list : ('k, 'v) t -> ('k * 'v) list + + val length : ('k, 'v) t -> int + (** Linear time: counts the bindings of a snapshot, like {!fold} *) + + val clear : ('k, 'v) t -> unit + (** Not atomic across shards *) +end + val locked : Mutex.t -> (unit -> 'a) -> 'a type 'a t From 3d31448a2e4169160e43a4328243cdd7966f28e4 Mon Sep 17 00:00:00 2001 From: Simon Cruanes Date: Wed, 30 Sep 2026 22:37:42 -0400 Subject: [PATCH 3/5] bench and tests for the new map --- bench/bench_sharded_hash_trie.ml | 159 +++++++++++++++++++++++++++++++ bench/dune | 6 ++ devkit.opam | 1 + tests/cache_count_test.ml | 27 ++++++ tests/dune | 3 + tests/sharded_hash_trie_test.ml | 150 +++++++++++++++++++++++++++++ 6 files changed, 346 insertions(+) create mode 100644 bench/bench_sharded_hash_trie.ml create mode 100644 bench/dune create mode 100644 tests/cache_count_test.ml create mode 100644 tests/dune create mode 100644 tests/sharded_hash_trie_test.ml diff --git a/bench/bench_sharded_hash_trie.ml b/bench/bench_sharded_hash_trie.ml new file mode 100644 index 00000000..7092010a --- /dev/null +++ b/bench/bench_sharded_hash_trie.ml @@ -0,0 +1,159 @@ +(* Compare ExtThread.ShardedHashTrie (64 and 256 shards) with Saturn.Htbl, a flat bucketed assoc list, + Hashtbl + Mutex, and (on one domain) a plain Hashtbl. + + dune exec --release bench/bench_sharded_hash_trie.exe *) + +open Devkit + +type ('k, 'v) ops = { find : 'k -> 'v; get_or_create : 'k -> f:('k -> 'v) -> 'v } +type 'k impl = { name : string; make : unit -> ('k, int Atomic.t) ops } + +module Impls (K : Hashtbl.HashedType) = struct + let devkit ?(shard_bits = 6) () = { name = Printf.sprintf "devkit/s%d" shard_bits; make = fun () -> + let t = ExtThread.ShardedHashTrie.create ~shard_bits ~hashable:{ equal = K.equal; hash = K.hash } () in + { find = ExtThread.ShardedHashTrie.find t; get_or_create = (fun k ~f -> ExtThread.ShardedHashTrie.get_or_create t k ~f) } } + + let saturn = { name = "saturn"; make = fun () -> + let t = Saturn.Htbl.create ~hashed_type:(module K) () in + let rec get_or_create k ~f = + match Saturn.Htbl.find_exn t k with + | v -> v + | exception Not_found -> let v = f k in if Saturn.Htbl.try_add t k v then v else get_or_create k ~f + in + { find = Saturn.Htbl.find_exn t; get_or_create } } + + (* fixed buckets of assoc lists, the simplest lock-free option *) + let flat = { name = "flat128"; make = fun () -> + let b = Array.init 128 (fun _ -> Atomic.make []) in + let bucket k = b.(K.hash k land 127) in + let rec assoc k = function [] -> raise_notrace Not_found | (k', v) :: tl -> if K.equal k k' then v else assoc k tl in + let find k = assoc k (Atomic.get (bucket k)) in + let rec get_or_create k ~f = + let b = bucket k in + let l = Atomic.get b in + match assoc k l with + | v -> v + | exception Not_found -> let v = f k in if Atomic.compare_and_set b l ((k, v) :: l) then v else get_or_create k ~f + in + { find; get_or_create } } + + let mutex = { name = "hashtbl+mutex"; make = fun () -> + let module T = Hashtbl.Make (K) in + let t = T.create 16 and m = Mutex.create () in + { find = (fun k -> Mutex.protect m (fun () -> T.find t k)); + get_or_create = (fun k ~f -> Mutex.protect m (fun () -> + match T.find t k with v -> v | exception Not_found -> let v = f k in T.add t k v; v)) } } + + (* not domain-safe: baseline for 1-domain runs only *) + let plain = { name = "hashtbl (unsafe)"; make = fun () -> + let module T = Hashtbl.Make (K) in + let t = T.create 16 in + { find = T.find t; + get_or_create = (fun k ~f -> match T.find t k with v -> v | exception Not_found -> let v = f k in T.add t k v; v) } } + + let tries = [ devkit (); devkit ~shard_bits:8 () ] + let all = tries @ [ saturn; flat; mutex ] + let growing = tries @ [ saturn; mutex ] (* flat128 degrades to long lists *) +end + +module S = Impls (struct include String let hash = Hashtbl.hash end) +module I = Impls (struct type t = int let equal = Int.equal let hash = Hashtbl.hash end) + +let counter _ = Atomic.make 0 + +let on_domains n f = List.init n (fun d -> Domain.spawn (fun () -> f d)) |> List.iter Domain.join + +let report ~title ~ops_per_call samples = + Printf.printf "\n== %s\n" title; + List.iter (fun (name, ts) -> + let rates = List.map (fun (t : Benchmark.t) -> Int64.to_float t.iters *. float ops_per_call /. t.wall /. 1e6) ts in + let rates = List.sort compare rates in + Printf.printf " %-16s %8.1f Mops/s (median of %d)\n%!" name (List.nth rates (List.length rates / 2)) (List.length rates)) + samples + +(* the unsynchronized Hashtbl only makes sense on one domain *) +let with_plain domains impls plain = if domains = 1 then impls @ [ plain ] else impls + +let bench ~title ~ops_per_call impls run = + let samples = Benchmark.throughputN ~style:Benchmark.Nil ~repeat:5 1 + (List.map (fun impl -> impl.name, run, impl) impls) in + report ~title ~ops_per_call samples + +(* Cache.Count: a few constant keys, hit on every call *) +let count_like domains = + let keys = Array.init 20 (Printf.sprintf "metric_name_%d") in + let per_domain = 1_000_000 in + bench ~title:(Printf.sprintf "count-like: 20 string keys, hits + incr, %d domain(s)" domains) + ~ops_per_call:(per_domain * domains) (with_plain domains S.all S.plain) + (fun impl -> + let t = impl.make () in + Array.iter (fun k -> ignore (t.get_or_create k ~f:counter)) keys; + on_domains domains (fun d -> + for i = 0 to per_domain - 1 do + Atomic.incr (t.get_or_create keys.((i + d) mod 20) ~f:counter) + done)) + +(* growing table: every call is a miss and an insert *) +let inserts domains = + let n = 200_000 in + bench ~title:(Printf.sprintf "inserts: %d fresh int keys, %d domain(s)" n domains) + ~ops_per_call:n (with_plain domains I.all I.plain) + (fun impl -> + let t = impl.make () in + on_domains domains (fun d -> + let i = ref d in + while !i < n do ignore (t.get_or_create !i ~f:counter); i := !i + domains done)) + +(* read-mostly: 10k warm keys, one insert of a new key every [insert_every] ops, finds otherwise. + The table grows by [per_domain / insert_every] keys per domain and call. *) +let mixed ~label ~insert_every domains = + let warm = 10_000 and per_domain = 2_000_000 in + bench ~title:(Printf.sprintf "%s: %d warm int keys, 1 insert per %d ops, %d domain(s)" label warm insert_every domains) + ~ops_per_call:(per_domain * domains) (with_plain domains I.growing I.plain) + (fun impl -> + let t = impl.make () in + for k = 0 to warm - 1 do ignore (t.get_or_create k ~f:counter) done; + on_domains domains (fun d -> + let fresh = ref (warm + d) in + for i = 0 to per_domain - 1 do + if i mod insert_every = 0 then (ignore (t.get_or_create !fresh ~f:counter); fresh := !fresh + domains) + else ignore (t.find ((i * 7919) mod warm)) + done)) + +let mixes = [ "mixed90", 10; "mixed99", 100; "mixed99.9", 1000 ] +let all_mixed () = List.iter (fun (label, insert_every) -> List.iter (mixed ~label ~insert_every) [ 1; 4; 8 ]) mixes + +(* sanity check every implementation before timing it *) +let self_check () = + List.iter (fun impl -> + let t = impl.make () in + let n = 100_000 in + on_domains 4 (fun d -> let i = ref d in while !i < n do ignore (t.get_or_create !i ~f:(fun k -> Atomic.make k)); i := !i + 4 done); + for k = 0 to n - 1 do if Atomic.get (t.find k) <> k then failwith (impl.name ^ ": self-check failed") done) + I.all + +let alloc_per_hit () = + Printf.printf "\n== minor words per hit (20 string keys, 1 domain)\n"; + let keys = Array.init 20 (Printf.sprintf "metric_name_%d") in + List.iter (fun impl -> + let t = impl.make () in + Array.iter (fun k -> ignore (t.get_or_create k ~f:counter)) keys; + let n = 1_000_000 in + let before = Gc.minor_words () in + for i = 0 to n - 1 do ignore (t.get_or_create keys.(i mod 20) ~f:counter) done; + Printf.printf " %-16s %6.2f\n%!" impl.name ((Gc.minor_words () -. before) /. float n)) + S.all + +let () = + self_check (); + match Sys.argv with + | [| _; "mixed" |] -> all_mixed () + | [| _; "single" |] -> + count_like 1; + inserts 1; + List.iter (fun (label, insert_every) -> mixed ~label ~insert_every 1) mixes + | _ -> + alloc_per_hit (); + List.iter count_like [ 1; 4; 8 ]; + List.iter inserts [ 1; 4 ]; + all_mixed () diff --git a/bench/dune b/bench/dune new file mode 100644 index 00000000..00c48b47 --- /dev/null +++ b/bench/dune @@ -0,0 +1,6 @@ +; optional: needs saturn and benchmark. Build explicitly: +; dune build --release bench/bench_sharded_hash_trie.exe +(executable + (name bench_sharded_hash_trie) + (optional) + (libraries devkit saturn benchmark unix)) diff --git a/devkit.opam b/devkit.opam index 31a089a3..519cef9d 100644 --- a/devkit.opam +++ b/devkit.opam @@ -15,6 +15,7 @@ depends: [ "dune" {>= "2.0"} ("extlib" {>= "1.7.1"} | "extlib-compat" {>= "1.7.1"}) "ounit2" + "qcheck-core" {with-test} "camlzip" "libevent" {>= "0.8.0"} "curl" {>= "0.10.0"} diff --git a/tests/cache_count_test.ml b/tests/cache_count_test.ml new file mode 100644 index 00000000..840ad578 --- /dev/null +++ b/tests/cache_count_test.ml @@ -0,0 +1,27 @@ +open Devkit +open QCheck2 +module C = Cache.Count + +(* each domain adds every key of [keys] once *) +let concurrent_add = + Test.make ~name:"concurrent add" ~count:50 + Gen.(pair (int_range 1 8) (list_size (int_range 0 20) (string_size (int_range 0 3)))) + (fun (n_domains, keys) -> + let c = C.create () in + List.init n_domains (fun _ -> Domain.spawn (fun () -> List.iter (C.add c) keys)) + |> List.iter Domain.join; + let expected = Hashtbl.create 16 in + List.iter (fun k -> Hashtbl.replace expected k (n_domains + Option.value ~default:0 (Hashtbl.find_opt expected k))) keys; + C.size c = Hashtbl.length expected + && Hashtbl.fold (fun k n ok -> ok && C.count c k = n) expected true + && C.count_all c = n_domains * List.length keys) + +let nan_key = + Test.make ~name:"nan key" ~count:1 Gen.unit (fun () -> + let c = C.create () in + C.add c nan; C.add c nan; + C.size c = 1 && C.count c nan = 2) + +let () = + ignore (Unix.alarm 300 : int); + exit (QCheck_base_runner.run_tests ~verbose:true [ concurrent_add; nan_key ]) diff --git a/tests/dune b/tests/dune new file mode 100644 index 00000000..df491fdf --- /dev/null +++ b/tests/dune @@ -0,0 +1,3 @@ +(tests + (names sharded_hash_trie_test cache_count_test) + (libraries devkit qcheck-core qcheck-core.runner unix)) diff --git a/tests/sharded_hash_trie_test.ml b/tests/sharded_hash_trie_test.ml new file mode 100644 index 00000000..48fea3ea --- /dev/null +++ b/tests/sharded_hash_trie_test.ml @@ -0,0 +1,150 @@ +open Devkit +module H = ExtThread.ShardedHashTrie + +(* hash functions to exercise the trie: collisions, deep splits, sign and high bits *) +let hashables = [ + "poly", H.poly_hashable; + "low byte", { H.equal = Int.equal; hash = (fun k -> k land 0xff) }; + "constant", { H.equal = Int.equal; hash = (fun _ -> 0) }; + "negative", { H.equal = Int.equal; hash = (fun k -> - (Hashtbl.hash k) - 1) }; + "high bits", { H.equal = Int.equal; hash = (fun k -> Hashtbl.hash k lsl 32) }; +] + +type op = Get_or_create of int * int | Find of int | Mem of int | Clear + +let show_op = function + | Get_or_create (k, v) -> Printf.sprintf "get_or_create %d %d" k v + | Find k -> Printf.sprintf "find %d" k + | Mem k -> Printf.sprintf "mem %d" k + | Clear -> "clear" + +let gen_ops = + let open QCheck2.Gen in + let* range = oneof_list [ 8; 1_000; 1_000_000 ] in + let key = int_bound range in + list_size (int_bound 2_000) @@ oneof_weighted [ + 10, map2 (fun k v -> Get_or_create (k, v)) key nat; + 10, map (fun k -> Find k) key; + 3, map (fun k -> Mem k) key; + 1, pure Clear; + ] + +let sorted l = List.sort compare l + +(* single domain: behaves like Hashtbl with add-if-absent *) +let model_test (name, hashable) = + QCheck2.Test.make ~count:300 ~name:("model vs Hashtbl, " ^ name) + ~print:QCheck2.Print.(list show_op) gen_ops + (fun ops -> + let t = H.create ~hashable ~shard_bits:2 () in + let m = Hashtbl.create 16 in + let step op = + match op with + | Get_or_create (k, v) -> + let expected = match Hashtbl.find_opt m k with Some v -> v | None -> Hashtbl.add m k v; v in + H.get_or_create t k ~f:(fun _ -> v) = expected + | Find k -> H.find_opt t k = Hashtbl.find_opt m k + | Mem k -> H.mem t k = Hashtbl.mem m k + | Clear -> H.clear t; Hashtbl.reset m; true + in + List.for_all step ops + && H.length t = Hashtbl.length m + && sorted (H.to_list t) = sorted (List.of_seq (Hashtbl.to_seq m))) + +let shard_bits_test = + QCheck2.Test.make ~count:50 ~name:"any shard_bits" + QCheck2.Gen.(pair (int_bound 16) (list_size (int_bound 500) nat)) + (fun (shard_bits, keys) -> + let t = H.create ~shard_bits () in + List.iter (fun k -> ignore (H.get_or_create t k ~f:Fun.id : int)) keys; + List.for_all (fun k -> H.find t k = k) keys + && H.length t = List.length (List.sort_uniq compare keys)) + +let invalid_shard_bits_test = + QCheck2.Test.make ~count:1 ~name:"shard_bits out of range" QCheck2.Gen.unit (fun () -> + List.for_all (fun shard_bits -> + match H.create ~shard_bits () with _ -> false | exception Invalid_argument _ -> true) + [ -1; 17 ]) + +(* Hashtbl.hash only looks at the first 10 meaningful values: these all collide *) +let structural_collision_test = + QCheck2.Test.make ~count:50 ~name:"full hash collisions (long lists)" + QCheck2.Gen.(list_size (int_range 1 200) nat) + (fun tails -> + let prefix = List.init 10 Fun.id in + let keys = List.sort_uniq compare tails |> List.map (fun x -> prefix @ [ x ]) in + let h0 = Hashtbl.hash (List.hd keys) in + assert (List.for_all (fun k -> Hashtbl.hash k = h0) keys); + let t = H.create () in + List.iteri (fun i k -> ignore (H.get_or_create t k ~f:(fun _ -> i) : int)) keys; + List.for_all2 (fun i k -> H.find t k = i) (List.init (List.length keys) Fun.id) keys + && H.length t = List.length keys) + +let long_strings_test = + QCheck2.Test.make ~count:50 ~name:"long string keys" + QCheck2.Gen.(list_size (int_bound 300) (string_size (int_range 100 2_000))) + (fun keys -> + let t = H.create ~hashable:{ H.equal = String.equal; hash = Hashtbl.hash } () in + List.iter (fun k -> ignore (H.get_or_create t k ~f:String.length : int)) keys; + List.for_all (fun k -> H.find t k = String.length k) keys + && H.length t = List.length (List.sort_uniq compare keys)) + +let large_test (name, hashable) = + let n = if name = "constant" then 2_000 else 200_000 in + QCheck2.Test.make ~count:1 ~name:(Printf.sprintf "%d keys, %s" n name) QCheck2.Gen.unit (fun () -> + let t = H.create ~hashable () in + for k = 0 to n - 1 do ignore (H.get_or_create t k ~f:(fun k -> 2 * k) : int) done; + let ok = ref (H.length t = n) in + for k = 0 to n - 1 do if H.find t k <> 2 * k then ok := false done; + !ok && not (H.mem t n) && H.fold t (fun _ _ c -> c + 1) 0 = n) + +let domains = 4 + +(* all domains race on the same keys: each key must end up with a single value, + and every caller must get that same (physical) value *) +let concurrent_get_or_create_test (name, hashable) = + QCheck2.Test.make ~count:20 ~name:("concurrent get_or_create, " ^ name) + QCheck2.Gen.(int_range 1 (if name = "constant" then 300 else 5_000)) + (fun n -> + let t = H.create ~hashable ~shard_bits:1 () in + let results = Array.init domains (fun d -> + Domain.spawn (fun () -> + (* different orders per domain to create contention everywhere *) + Array.init n (fun i -> + let k = if d land 1 = 0 then i else n - 1 - i in + k, H.get_or_create t k ~f:(fun k -> ref k)))) + |> Array.map Domain.join + in + let per_key = Array.make n [] in + Array.iter (Array.iter (fun (k, v) -> per_key.(k) <- v :: per_key.(k))) results; + H.length t = n + && Array.for_all (fun vs -> let v = H.find t !(List.hd vs) in List.for_all (fun v' -> v' == v) vs) per_key) + +(* a key found once is found forever, while another domain keeps inserting *) +let concurrent_readers_test = + QCheck2.Test.make ~count:10 ~name:"readers during writes" QCheck2.Gen.(int_range 1_000 50_000) + (fun n -> + let t = H.create ~shard_bits:2 () in + let writer = Domain.spawn (fun () -> for k = 0 to n - 1 do ignore (H.get_or_create t k ~f:Fun.id : int) done) in + let readers = List.init (domains - 1) (fun _ -> Domain.spawn (fun () -> + let seen = ref 0 and ok = ref true in + while !seen < n do + (* everything below [seen] was found before, must still be there *) + for k = max 0 (!seen - 100) to !seen - 1 do if H.find_opt t k <> Some k then ok := false done; + if H.mem t !seen then incr seen else Domain.cpu_relax () + done; + !ok)) in + Domain.join writer; + List.for_all Domain.join readers && H.length t = n) + +let () = + (* a broken trie tends to loop forever rather than fail: the default SIGALRM action kills us *) + ignore (Unix.alarm 300 : int); + let tests = + List.map model_test hashables + @ [ shard_bits_test; invalid_shard_bits_test; structural_collision_test; long_strings_test ] + @ List.map large_test hashables + @ List.map concurrent_get_or_create_test hashables + @ [ concurrent_readers_test ] + in + exit (QCheck_base_runner.run_tests ~verbose:true tests) From 69d5fed60bebfcef76546874d9c8a3c21ae8c047 Mon Sep 17 00:00:00 2001 From: Simon Cruanes Date: Thu, 1 Oct 2026 11:12:16 -0400 Subject: [PATCH 4/5] make Cache.Count domain-safe --- cache.ml | 46 ++++++++++++++++++++++++---------------------- cache.mli | 7 ++++++- 2 files changed, 30 insertions(+), 23 deletions(-) diff --git a/cache.ml b/cache.ml index 830527ef..222458e2 100644 --- a/cache.ml +++ b/cache.ml @@ -60,37 +60,39 @@ module TimeLimited2(E: Set.OrderedType) end module Count = struct - open Hashtbl - type 'a t = ('a, int ref) Hashtbl.t - let create () : 'a t = create 16 - let clear = Hashtbl.clear - let entry t x = match find t x with r -> r | exception Not_found -> let r = ref 0 in Hashtbl.add t x r; r - let plus t x n = entry t x += n - let minus t x n = entry t x -= n + module H = ExtThread.ShardedHashTrie + (* counters are atomic, the table only grows (or is cleared) *) + type 'a t = ('a, int Atomic.t) H.t + let create () : 'a t = H.create ~shard_bits:6 () + let clear = H.clear + let entry t x = H.get_or_create t x ~f:(fun _ -> Atomic.make 0) + let plus t x n = ignore (Atomic.fetch_and_add (entry t x) n : int) + let minus t x n = plus t x (-n) let of_enum e = let h = create () in Enum.iter (fun (k,n) -> plus h k n) e; h let of_list l = of_enum @@ List.enum l let add t x = plus t x 1 let del t x = minus t x 1 - let enum t = enum t |> Enum.map (fun (k,n) -> k, !n) - let iter t f = iter (fun k n -> f k !n) t - let fold t f acc = Hashtbl.fold (fun k n acc -> f k !n acc) t acc - let count t k = match Hashtbl.find t k with n -> !n | exception Not_found -> 0 - let count_all t = Hashtbl.fold (fun _ n acc -> acc + !n) t 0 - let size = Hashtbl.length - let show t ?(sep=" ") f = enum t |> - List.of_enum |> List.sort ~cmp:(Action.compare_by fst) |> + let iter t f = H.iter t (fun k n -> f k (Atomic.get n)) + let fold t f acc = H.fold t (fun k n acc -> f k (Atomic.get n) acc) acc + let to_list t = fold t (fun k n acc -> (k,n) :: acc) [] + let enum t = List.enum (to_list t) + let count t k = match H.find_opt t k with Some n -> Atomic.get n | None -> 0 + let count_all t = fold t (fun _ n acc -> acc + n) 0 + let size = H.length + let show t ?(sep=" ") f = to_list t |> + List.sort ~cmp:(Action.compare_by fst) |> List.map (fun (x,n) -> sprintf "%S: %u" (f x) n) |> String.concat sep - let show_sorted t ?limit ?(sep="\n") f = enum t |> - List.of_enum |> List.sort ~cmp:(flip @@ Action.compare_by snd) |> + let show_sorted t ?limit ?(sep="\n") f = to_list t |> + List.sort ~cmp:(flip @@ Action.compare_by snd) |> (match limit with None -> id | Some n -> List.take n) |> List.map (fun (x,n) -> sprintf "%6d : %S" n (f x)) |> String.concat sep let stats t ?(cmp=compare) f = - if Hashtbl.length t = 0 then + let a = to_list t |> Array.of_list in + if Array.length a = 0 then "" else - let a = Array.of_enum (enum t) in let total = Array.fold_left (fun t (_,n) -> t + n) 0 a in let half = total / 2 in let cmp (x,_) (y,_) = cmp x y in @@ -108,10 +110,10 @@ module Count = struct sprintf "total %d median %s min %s max %s" total (match !med with None -> "?" | Some x -> show x) (show mi) (show ma) let distrib t = - if Hashtbl.length t = 0 then + let a = to_list t |> Array.of_list in + if Array.length a = 0 then [||] else - let a = Array.of_enum (enum t) in let total = Array.fold_left (fun t (_,n) -> t + n) 0 a in let limits = Array.init 10 (fun i -> total * (i + 1) / 10) in let cmp (x,_) (y,_) = compare (x:float) y in @@ -133,7 +135,7 @@ module Count = struct let data = show_sorted t ?limit ~sep f in let stats = stats t ?cmp f in stats^sep^data - let names (t : 'a t) = List.of_enum @@ Hashtbl.keys t + let names (t : 'a t) = H.fold t (fun k _ acc -> k :: acc) [] end diff --git a/cache.mli b/cache.mli index e6c5e40f..f6b71c92 100644 --- a/cache.mli +++ b/cache.mli @@ -38,7 +38,12 @@ module LRU(K : Hashtbl.HashedType) : sig val lfu_free : 'v t -> int end -(** Count elements *) +(** Count elements. Domain safe. + + Collection-wide operations such as [iter], [fold], [size], + [clear], etc. can race against individual modifications (ie it's + possible that a call to [clear t] races against [plus t "x" 42] + and an empty structure is never observable) *) module Count : sig type 'a t val create : unit -> 'a t From 7dc1b13884b20000dce8deecff7edfa2a74e9410 Mon Sep 17 00:00:00 2001 From: Simon Cruanes Date: Thu, 1 Oct 2026 13:52:14 -0400 Subject: [PATCH 5/5] make benchmark a dev-tool only dep --- Makefile | 5 ++++- bench/dune | 4 ++-- devkit.opam | 3 ++- dune-project | 2 +- 4 files changed, 9 insertions(+), 5 deletions(-) diff --git a/Makefile b/Makefile index 2c8fd7d2..0e6af992 100644 --- a/Makefile +++ b/Makefile @@ -1,5 +1,5 @@ -.PHONY: build lib doc clean install uninstall test gen gen_ragel gen_metaocaml archive +.PHONY: build lib doc clean install uninstall test benchs gen gen_ragel gen_metaocaml archive OCAMLBUILD=ocamlbuild -use-ocamlfind -no-links -j 0 @@ -27,6 +27,9 @@ top: test: dune runtest $(DUNEFLAGS) +benchs: + RUN_BENCHS=1 dune exec --profile bench $(DUNEFLAGS) bench/bench_sharded_hash_trie.exe + doc: dune build $(DUNEFLAGS) @doc diff --git a/bench/dune b/bench/dune index 00c48b47..c2e7e1fe 100644 --- a/bench/dune +++ b/bench/dune @@ -1,6 +1,6 @@ -; optional: needs saturn and benchmark. Build explicitly: -; dune build --release bench/bench_sharded_hash_trie.exe +; optional: needs saturn and benchmark. Build with make benchs. (executable (name bench_sharded_hash_trie) (optional) + (enabled_if (= %{profile} bench)) (libraries devkit saturn benchmark unix)) diff --git a/devkit.opam b/devkit.opam index 519cef9d..3940a5fd 100644 --- a/devkit.opam +++ b/devkit.opam @@ -12,10 +12,11 @@ build: [ ] depends: [ "ocaml" {>= "5.0"} - "dune" {>= "2.0"} + "dune" {>= "2.3"} ("extlib" {>= "1.7.1"} | "extlib-compat" {>= "1.7.1"}) "ounit2" "qcheck-core" {with-test} + "benchmark" {with-dev-setup} "camlzip" "libevent" {>= "0.8.0"} "curl" {>= "0.10.0"} diff --git a/dune-project b/dune-project index 3411e7fe..4c75eda4 100644 --- a/dune-project +++ b/dune-project @@ -1,3 +1,3 @@ -(lang dune 2.0) +(lang dune 2.3) (name devkit) (implicit_transitive_deps false)