Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 8 additions & 10 deletions common/Switch.v
Original file line number Diff line number Diff line change
Expand Up @@ -18,7 +18,7 @@
to comparison trees. *)

From Coq Require Import EqNat.
Require Import Coqlib Maps Integers Values.
Require Import Coqlib Zbits Maps Integers Values.

(** A multi-way branch is composed of a list of (key, action) pairs,
plus a default action. *)
Expand Down Expand Up @@ -121,7 +121,7 @@ Fixpoint validate_jumptable (cases: ZMap.t nat)
match tbl with
| nil => true
| act :: rem =>
Nat.eqb act (ZMap.get n cases)
Nat.eqb act (ZMap.get (n mod modulus) cases)
&& validate_jumptable cases rem (Z.succ n)
end.

Expand Down Expand Up @@ -157,7 +157,7 @@ Fixpoint validate (default: nat) (cases: table) (t: comptree)
let tbl_len := list_length_z tbl in
zle 0 ofs && zlt ofs modulus &&
zle 0 sz && zlt sz modulus &&
zle (ofs + sz) modulus && zle sz tbl_len && zlt sz Int.modulus &&
zle sz tbl_len && zlt sz Int.modulus &&
match split_between default ofs sz cases with
| (inside, outside) =>
validate_jumptable inside tbl ofs
Expand Down Expand Up @@ -264,7 +264,7 @@ Lemma validate_jumptable_correct_rec:
forall cases tbl base v,
validate_jumptable cases tbl base = true ->
0 <= v < list_length_z tbl ->
list_nth_z tbl v = Some(ZMap.get (base + v) cases).
list_nth_z tbl v = Some(ZMap.get ((base + v) mod modulus) cases).
Proof.
induction tbl; simpl; intros.
- unfold list_length_z in H0. simpl in H0. extlia.
Expand All @@ -279,18 +279,16 @@ Lemma validate_jumptable_correct:
forall cases tbl ofs v sz,
validate_jumptable cases tbl ofs = true ->
(v - ofs) mod modulus < sz ->
0 <= sz -> 0 <= ofs -> ofs + sz <= modulus ->
0 <= v < modulus ->
sz <= list_length_z tbl ->
list_nth_z tbl ((v - ofs) mod modulus) = Some(ZMap.get v cases).
Proof.
intros.
rewrite (validate_jumptable_correct_rec cases tbl ofs); auto.
- f_equal. f_equal. rewrite Z.mod_small. lia.
destruct (zle ofs v). lia.
assert (M: ((v - ofs) + 1 * modulus) mod modulus = (v - ofs) + modulus).
{ rewrite Z.mod_small. lia. lia. }
rewrite Z_mod_plus in M by auto. rewrite M in H0. lia.
- f_equal. f_equal. rewrite <- (Z.mod_small v modulus) at 2 by lia.
apply eqmod_mod_eq; auto.
replace v with (ofs + (v - ofs)) at 2 by lia.
auto using eqmod_add, eqmod_sym, eqmod_mod, eqmod_refl.
- generalize (Z_mod_lt (v - ofs) modulus modulus_pos). lia.
Qed.

Expand Down
39 changes: 25 additions & 14 deletions common/Switchaux.ml
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ module ZSet = Set.Make(Z)

let normalize_table tbl =
let rec norm keys accu = function
| [] -> (accu, keys)
| [] -> accu
| (key, act) :: rem ->
if ZSet.mem key keys
then norm keys accu rem
Expand Down Expand Up @@ -66,17 +66,18 @@ let compile_switch_as_tree modulus default tbl =
build mid hi pivot maxval)
in build 0 (Array.length sw) Z.zero modulus

let compile_switch_as_jumptable default cases minkey maxkey =
let tblsize = 1 + Z.to_int (Z.sub maxkey minkey) in
assert (tblsize >= 0 && tblsize <= Sys.max_array_length);
let tbl = Array.make tblsize default in
let compile_switch_as_jumptable modulus default cases minkey maxkey =
let size = Z.(add (sub maxkey minkey) one) in
assert (Z.gt size Z.zero && Z.le size (Z.of_uint Sys.max_array_length));
let tbl = Array.make (Z.to_int size) default in
List.iter
(fun (key, act) ->
let pos = Z.to_int (Z.sub key minkey) in
tbl.(pos) <- act)
let pos = Z.(modulo (sub key minkey) modulus) in
assert (Z.ge pos Z.zero && Z.lt pos size);
tbl.(Z.to_int pos) <- act)
cases;
CTjumptable(minkey,
Z.of_uint tblsize,
CTjumptable(Z.modulo minkey modulus,
size,
Array.to_list tbl,
CTaction default)

Expand All @@ -89,13 +90,23 @@ let dense_enough (numcases: int) (minkey: Z.t) (maxkey: Z.t) =
&& Z.le table_size tree_size
&& Z.lt span (Z.of_uint Sys.max_array_length)

let signed_min_max_key modulus tbl =
let half_modulus = Z.(shr modulus 1) in
let signed n =
if Z.lt n half_modulus then n else Z.sub n modulus in
let rec min_max lo hi = function
| [] -> (lo, hi)
| (key, _) :: tbl ->
let skey = signed key in
min_max (Z.min skey lo) (Z.max skey hi) tbl in
min_max (Z.pred half_modulus) (Z.neg half_modulus) tbl

let compile_switch modulus default table =
let (tbl, keys) = normalize_table table in
if ZSet.is_empty keys then CTaction default else begin
let minkey = ZSet.min_elt keys
and maxkey = ZSet.max_elt keys in
let tbl = normalize_table table in
if tbl = [] then CTaction default else begin
let (minkey, maxkey) = signed_min_max_key modulus tbl in
if dense_enough (List.length tbl) minkey maxkey
then compile_switch_as_jumptable default tbl minkey maxkey
then compile_switch_as_jumptable modulus default tbl minkey maxkey
else compile_switch_as_tree modulus default tbl
end

Expand Down
2 changes: 1 addition & 1 deletion test
Loading