summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorBrian Campbell2017-11-27 16:41:58 +0000
committerBrian Campbell2017-11-27 16:42:55 +0000
commit4a7d6e6d7e9221a19bc50c627b5714e45b1748bc (patch)
treec22a44c80e45b17c1c96c802450aa8eeb2bc1fe7
parent05081c479140487dd581167168f67e8226ec599f (diff)
Use guards from when patterns when typing cases
-rw-r--r--src/type_check.ml4
-rw-r--r--test/typecheck/pass/patternrefinement.sail19
2 files changed, 22 insertions, 1 deletions
diff --git a/src/type_check.ml b/src/type_check.ml
index 3a8f2e59..2260416d 100644
--- a/src/type_check.ml
+++ b/src/type_check.ml
@@ -2041,7 +2041,9 @@ let rec check_exp env (E_aux (exp_aux, (l, ())) as exp : unit exp) (Typ_aux (typ
| Pat_aux (Pat_when (pat, guard, case), (l, _)) ->
let tpat, env = bind_pat env pat (typ_of inferred_exp) in
let checked_guard = check_exp env guard bool_typ in
- Pat_aux (Pat_when (tpat, checked_guard, crule check_exp env case typ), (l, None))
+ let flows, constrs = infer_flow env checked_guard in
+ let env' = add_constraints constrs (add_flows true flows env) in
+ Pat_aux (Pat_when (tpat, checked_guard, crule check_exp env' case typ), (l, None))
in
annot_exp (E_case (inferred_exp, List.map (fun case -> check_case case typ) cases)) typ
| E_try (exp, cases), _ ->
diff --git a/test/typecheck/pass/patternrefinement.sail b/test/typecheck/pass/patternrefinement.sail
new file mode 100644
index 00000000..5a89372a
--- /dev/null
+++ b/test/typecheck/pass/patternrefinement.sail
@@ -0,0 +1,19 @@
+default Order dec
+
+val extern forall Num 'n, Num 'm, Num 'o, Num 'p, Order 'ord.
+ vector<'o, 'n, 'ord, bit> -> vector<'p, 'm, 'ord, bit> effect pure extz
+val extern forall Num 'n, Num 'm, Order 'ord, Type 'a. vector<'n,'m,'ord,'a> -> [:'m:] effect pure length
+val extern forall Num 'n, Num 'm, Order 'ord. (vector<'n,'m,'ord,bit>, vector<'n,'m,'ord,bit>) -> bool effect pure eq_vec
+val extern forall Num 'n, Num 'm. ([:'n:],[:'m:]) -> bool effect pure eq_atom
+val extern forall Type 'a. ('a, 'a) -> bool effect pure eq
+overload (deinfix ==) [eq_vec; eq_atom; eq]
+
+
+val forall 'n, 'n in {32,64}. bit['n] -> bit[64] effect pure test
+
+function test(v) = {
+ switch (length(v)) {
+ case 32 -> extz(v)
+ case 64 -> v
+ }
+}