File ‹Tools/smt_word.ML›
signature SMT_WORD =
sig
val mk_bv_from_int_base: term -> term -> term
end
structure SMT_Word : SMT_WORD =
struct
open Word_Lib
fun smtlib_logic "z3" ts _ =
if exists (Term.exists_type (Term.exists_subtype is_wordT)) ts
then SOME (SMT_Translate.NO_LOGIC)
else NONE
| smtlib_logic "verit" _ _ = NONE
| smtlib_logic _ ts _ =
if exists (Term.exists_type (Term.exists_subtype is_wordT)) ts
then SOME (SMT_Translate.SL ({fixedSizeBitVectors=true,real=true,set=false,
datatypes=false}:SMT_Translate.smtlib_theories))
else NONE
val smtlibC = SMTLIB_Interface.bvsmtlibC
val wordT = \<^typ>‹'a::len word›
fun index1 s i = "(_ " ^ s ^ " " ^ string_of_int i ^ ")"
fun index2 s i j = "(_ " ^ s ^ " " ^ string_of_int i ^ " " ^ string_of_int j ^ ")"
val mk_nat = HOLogic.mk_number \<^typ>‹nat›
fun remove_cast (Const (\<^const_name>‹Int.nat›, _) $ x) = x |
remove_cast x = x
fun if_fixed pred m n T ts =
let
val (Us, U) = Term.strip_type T
in
if pred (U, Us) then SOME (n, length Us, ts, Term.list_comb o pair (Const (m, T))) else NONE
end
fun if_concrete_bw_all_args m = if_fixed (forall (can dest_wordT) o (op ::)) m
fun if_fixed_args m = if_fixed (forall (can dest_wordT) o snd) m
fun add_word_fun f (t, n) =
let
val m = dest_Const_name t
in
SMT_Builtin.add_builtin_fun smtlibC (Term.dest_Const t, K (f m n))
end
fun word_typ (Type (\<^type_name>‹word›, [T])) = Option.map (rpair [] o index1 "BitVec") (try dest_binT T) |
word_typ (Type (_, [T])) = Option.map (rpair [] o index1 "BitVec") (try dest_binT T) |
word_typ _ = NONE
fun word_num (Type (\<^type_name>‹word›, [T])) k =
let
val size = try dest_binT T
fun max_int size = Integer.pow size 2
in
(case size of
NONE => NONE
| SOME size => SOME (index1 ("bv" ^ string_of_int (Int.rem(k, max_int size))) size))
end
| word_num _ _ = NONE
fun mk_shift c [u, t] = Const c $ mk_nat (snd (HOLogic.dest_number u)) $ t
| mk_shift c ts = raise TERM ("bad arguments", Const c :: ts)
fun shift m n T ts =
let val U = Term.domain_type (Term.range_type T)
in
(case (can dest_wordT U, try (snd o HOLogic.dest_number o hd) ts) of
(true, SOME i) =>
SOME (n, 2, [hd (tl ts), HOLogic.mk_number U i], mk_shift (m, T))
| _ => NONE)
end
fun mk_extract c i ts = Term.list_comb (Const c, mk_nat i :: ts)
fun slice m n (T as (Type(_,[_,Type(_,[Tx,U])]))) ts =
(case (try (snd o HOLogic.dest_number o remove_cast o hd) ts,
try dest_wordT U,
try dest_wordT Tx) of
(SOME j, SOME k, SOME i') =>
let
val i = i' - 1
in
if i >= j andalso k = i' - j
then (SOME (index2 n i j, 1, tl ts, mk_extract (m, T) j))
else NONE
end
| _ => NONE) |
slice _ _ _ _ = (NONE)
fun take_bit m _ T ts =
let
val ts = ts |> map (fn (Const (\<^const_name>‹Int.nat›, _) $ x) => x | x => x)
val word= hd (tl ts)
val U = Term.domain_type (Term.range_type T)
in
(case (can dest_wordT U, try (snd o HOLogic.dest_number o hd) ts) of
(true, SOME i) =>
let
val str=""
val offset' = HOLogic.mk_number @{typ "nat"} i
val minus= Const(\<^const_name>‹minus›, U --> U --> U)
val drop= Const(\<^const_name>‹drop_bit›, @{typ "nat"} --> U --> U)
val push= Const(\<^const_name>‹push_bit›, @{typ "nat"} --> U --> U)
fun mk_take_bit _ ts = hd ts
val new_term = minus $ word $
(push $ offset' $ (drop $ offset' $ word))
in
SOME (str, 1, [new_term], mk_take_bit (m, T)) end
| _ => NONE)
end
fun bit m n T ts =
let
val word = hd ts
val pos = hd (tl ts)
val suc_simplified =
(case (try (snd o HOLogic.dest_number o remove_cast) pos) of
SOME i' => (HOLogic.mk_number @{typ "nat"} (i'+1))
| NONE => (Const(\<^const_name>‹Nat.Suc›,@{typ "nat"} --> @{typ "nat"}) $ pos))
val pos_simplified =
(case (try (snd o HOLogic.dest_number o remove_cast) pos) of
SOME i' => (HOLogic.mk_number @{typ "nat"} (i'))
| NONE => (Const(\<^const_name>‹Nat.Suc›,@{typ "nat"} --> @{typ "nat"}) $ pos))
fun cond i i_suc =
(Const("smt_extract", @{typ "nat"} --> @{typ "nat"} --> fastype_of (word) -->
@{typ "1 word"}) $ i $ i $ word)
fun new_ts i i_suc =
Const(\<^const_name>‹HOL.eq›, @{typ "1 word"} --> @{typ "1 word"} --> @{typ "bool"}) $
cond i i_suc $
@{term "1::1 word"}
fun mk_bit ts = (hd ts)
in SOME ("", 1, [new_ts pos_simplified suc_simplified ], mk_bit) end
fun mk_extend c ts = Term.list_comb (Const c, ts)
fun extend m n T ts =
let val (U1, U2) = Term.dest_funT T
in
(case (try dest_wordT U1, try dest_wordT U2) of
(SOME i, SOME j) =>
if j-i >= 0 then SOME (index1 n (j-i), 1, ts, mk_extend (m, T))
else NONE
| _ => NONE)
end
fun mk_extend (m,T) ts = Term.list_comb (Const (m,T),ts)
fun is_numeral x = try HOLogic.dest_numeral x
fun power m n T ts =
let
val (U1, U2) = Term.dest_funT T
val bitwidth = U1 |> Term.dest_Type |> snd |> hd |> dest_binT
val wordT = mk_wordT bitwidth
val one = Const(\<^const_name>‹Groups.one_class.one›, wordT)
val t1 = nth ts 1
val (t2 $ t3) = t1
val shift_amount =
(case is_numeral t3 of
NONE => @{term "unsigned"} $ t3 |
SOME _ => Const (\<^const_name>‹Num.numeral_class.numeral›, @{typ "num"} --> wordT) $ t3)
in
(SOME ("bvshl", 2, [one,shift_amount], (fn _ => mk_extend (m,T) ts) ))
end
fun len_of_lift m n T [t1] =
let
val (_,b) = dest_Const t1
val (_,[c]) = dest_Type b
in
(case try dest_binT c of
NONE => NONE |
SOME len =>
let
val d = len |> HOLogic.mk_number @{typ int}
in
(SOME ("", 1, [d], fn _ => d))
end
) end |
len_of_lift _ _ _ _ = NONE
fun mk_rotate c i ts = Term.list_comb (Const c, mk_nat i :: ts)
fun rotate m n T ts =
let val U = Term.domain_type (Term.range_type T)
in
(case (can dest_wordT U, try (snd o HOLogic.dest_number o hd) ts) of
(true, SOME i) => SOME (index1 n i, 1, tl ts, mk_rotate (m, T) i)
| _ => NONE)
end
fun len_of t m n T ts =
let
fun mk_num T x =
let
val _ = HOLogic.mk_number HOLogic.intT (dest_binT T)
in Const (\<^const_name>‹nat›, @{typ "int => nat"}) $
HOLogic.mk_number HOLogic.intT (dest_binT T) end
in
(case head_of (hd ts) of
Const (_, Type (\<^type_name>‹itself›, [T])) =>
SOME ("len_of" , 0, [], (mk_num T))
| _ => NONE)
end
fun add_len_of (t,n) =
let val (m, _) = Term.dest_Const (head_of t)
in SMT_Builtin.add_builtin_fun smtlibC (Term.dest_Const (head_of t), fn _ => (fn x => fn y =>len_of t m n x y))
end
fun mk_op m T ts = Term.list_comb (Const(m,T),ts)
fun word_of_int m n T ts =
let
val U = Term.range_type T
in
(case try dest_wordT U of
NONE => NONE |
SOME U' => SOME (index1 n U', 1, ts,mk_op m T))
end
fun word_constructor m n T [t1] =
let
in
(SOME (n, 1, [t1], (fn t =>mk_extend (m,T) [t1])))
end |
word_constructor _ _ _ _ = NONE
fun of_nat t m n T ts =
let
fun mk_num T x =
let
val _ = HOLogic.mk_number HOLogic.natT (dest_binT T)
in HOLogic.mk_number HOLogic.natT (dest_binT T) end
in
(case (head_of (hd ts)) of
Const (_, Type (\<^type_name>‹itself›, [T])) =>
SOME ("LENGTHXXXXX", 0, [], (mk_num T))
| _ => NONE)
end
fun add_of_nat (a as t,n) =
let val (m, _) = Term.dest_Const t
val n = HOLogic.dest_number t |> snd |> K true
handle TERM _ => false
in
if n then SMT_Builtin.add_builtin_fun smtlibC (Term.dest_Const (head_of t), fn x => (@{print} x;
fn x => fn y =>(of_nat t m n x y)))
else SMT_Builtin.add_builtin_fun smtlibC (Term.dest_Const (head_of t), K (K (K (NONE))))
end
fun concat op_name smt_name (typ as (Type(_,[T1,Type(_,[T2,T])]))) terms =
(case (try dest_wordT T1, try dest_wordT T2, try dest_wordT T) of
(SOME i, SOME j, SOME k) =>
(if (i + j = k)
then (if_fixed (forall (can dest_wordT) o (op ::)) op_name smt_name typ terms)
else NONE) |
_ => NONE) |
concat _ _ _ _ = NONE
val bit_ops_table = [
(\<^const_name>‹bit›, @{thm min_def_raw})]
val setup_bit_ops_table = fold (SMT_Builtin.add_builtin_fun_ext'' o fst) bit_ops_table
val _ = Theory.setup (Context.theory_map (setup_bit_ops_table))
val setup_builtins =
SMT_Builtin.add_builtin_typ SMTLIB_Interface.bvsmtlibC (wordT, word_typ, word_num) #>
SMT_Builtin.add_builtin_typ SMTLIB_Interface.bvsmtlibC
(\<^typ>‹Num.num ⇒ 'a::len word›, word_typ, word_num) #>
fold (add_word_fun if_concrete_bw_all_args) [
(\<^term>‹uminus :: 'a::len word ⇒ _›, "bvneg"),
(\<^term>‹plus :: 'a::len word ⇒ _›, "bvadd"),
(\<^term>‹minus :: 'a::len word ⇒ _›, "bvsub"),
(\<^term>‹times :: 'a::len word ⇒ _›, "bvmul"),
(\<^term>‹not :: 'a::len word ⇒ _›, "bvnot"),
(\<^term>‹and :: 'a::len word ⇒ _›, "bvand"),
(\<^term>‹or :: 'a::len word ⇒ _›, "bvor"),
(\<^term>‹xor :: 'a::len word ⇒ _›, "bvxor"),
(\<^term>‹word_cat :: 'a::len word ⇒ _›, "concat") ] #>
fold (add_word_fun shift) [
(\<^term>‹push_bit :: nat ⇒ 'a::len word ⇒ _ ›, "bvshl"),
(\<^term>‹drop_bit :: nat ⇒ 'a::len word ⇒ _›, "bvlshr"),
(\<^term>‹signed_drop_bit :: nat ⇒ 'a::len word ⇒ _›, "bvashr") ] #>
add_word_fun concat
(\<^term>‹word_cat :: 'a::len word ⇒ _›, "concat") #>
add_word_fun slice
(\<^term>‹slice :: _ ⇒ 'a::len word ⇒ _›, "extract") #>
fold (add_word_fun extend) [
(\<^term>‹ucast :: 'a::len word ⇒ _›, "zero_extend"),
(\<^term>‹scast :: 'a::len word ⇒ _›, "sign_extend") ] #>
fold (add_word_fun rotate) [
(\<^term>‹word_rotl›, "rotate_left"),
(\<^term>‹word_rotr›, "rotate_right") ] #>
fold (add_word_fun if_fixed_args) [
(\<^term>‹less :: 'a::len word ⇒ _›, "bvult"),
(\<^term>‹less_eq :: 'a::len word ⇒ _›, "bvule"),
(\<^term>‹word_sless›, "bvslt"),
(\<^term>‹word_sle›, "bvsle") ] #>
fold (SMT_Builtin.add_builtin_fun' SMTLIB_Interface.bvsmtlibC) [
(\<^term>‹unsigned :: 'a::len word ⇒ Int.int›, "ubv_to_int"),
(\<^term>‹unsigned :: 'a::len word ⇒ Nat.nat›, "ubv_to_int")
] #>
fold (add_word_fun word_of_int) [
(\<^term>‹of_int :: Int.int => 'a::len word›, "int_to_bv"),
(\<^term>‹of_nat :: Nat.nat => 'a::len word›, "int_to_bv"),
(\<^term>‹Word.Word :: Int.int ⇒ 'a::len word›, "int_to_bv")
] #>
fold (add_word_fun take_bit) [
(\<^term>‹take_bit :: _ ⇒ 'a::len word ⇒ _›, "TODO")
] #>
fold (add_word_fun bit) [
(\<^term>‹bit :: 'a::len word ⇒ _ ⇒ _›, "bit")
]
open SMT_Parser_Util
val split_last =
let fun split_last [a] = ([], a)
| split_last (x :: xs) = apfst (curry (op ::) x) (split_last xs)
in apfst List.rev o split_last end
fun mk_word_of_int i u = SOME (Const (\<^const_name>‹of_int›, @{typ int} --> (Word_Lib.mk_wordT i)) $ u)
fun mk_word_of_int' i u = SOME (Const (\<^const_name>‹of_int›, @{typ int} --> (Word_Lib.mk_wordT (snd (HOLogic.dest_number i)))) $ u)
fun mk_zero_extend i u =
let
val T = fastype_of u
in
case (try dest_wordT T) of
(SOME j) => SOME (Const (\<^const_name>‹Word.cast›, T --> (Word_Lib.mk_wordT (i+j))) $ u) |
_ => NONE
end
fun mk_scast i u =
let
val T = fastype_of u
val j = Word_Lib.dest_wordT T
val TU = Word_Lib.mk_wordT (i + j)
in Const (\<^const_name>‹Word.signed›, T --> TU) $ u end;
fun mk_bv_from_int_base int base =
let
val ty = Word_Lib.mk_wordT (snd (HOLogic.dest_number base))
val num = snd (HOLogic.dest_number int)
handle TERM ("dest_number", [t ]) => (@{print}("t",t); raise Alethe_Proof.ALETHE_PROOF_PARSE "Could not parse bit-vector value")
in
(HOLogic.mk_number ty num)
end
fun
bv_term_parser (SMTLIB.BVNum (i, base), []) = SOME (HOLogic.mk_number (mk_wordT(base)) (i) )
| bv_term_parser (SMTLIB.Sym "bvnot", [t1]) =
SOME (mk_unary \<^const_name>‹ring_bit_operations_class.not› t1)
| bv_term_parser (SMTLIB.Sym "bvneg", [t]) =
SOME (mk_unary \<^const_name>‹uminus_class.uminus› t)
| bv_term_parser (SMTLIB.Sym "bvand", (t::ts)) =
SOME (mk_lassoc' \<^const_name>‹semiring_bit_operations_class.and› t ts)
| bv_term_parser (SMTLIB.Sym "bv", [int, base]) = SOME (mk_bv_from_int_base int base)
| bv_term_parser (SMTLIB.S [SMTLIB.Sym "_",SMTLIB.Sym n, SMTLIB.Num base],[]) =
(case try (unprefix "bv") n of
NONE => NONE
| SOME t =>
(case Int.fromString t of
NONE => NONE
|SOME t' =>
let
val int = HOLogic.mk_number @{typ int} t'
in
SOME (mk_bv_from_int_base int (HOLogic.mk_number @{typ int} base))
end)
)
| bv_term_parser (SMTLIB.Sym "_", [Free(n,_),base]) =
(case try (unprefix "bv") n of
NONE => NONE
| SOME t =>
(case Int.fromString t of
NONE => NONE
|SOME t' =>
let
val int = HOLogic.mk_number @{typ int} t'
in
SOME (mk_bv_from_int_base int base)
end)
)
| bv_term_parser (SMTLIB.S [SMTLIB.Sym "_", SMTLIB.Sym "int_to_bv", SMTLIB.Num i], [t]) =
mk_word_of_int i t
| bv_term_parser (SMTLIB.Sym "int_to_bv", [i,t]) =
mk_word_of_int' i t
| bv_term_parser (SMTLIB.Sym "ubv_to_int",[t1]) =
let
val T1 = fastype_of t1
in
SOME (Const (\<^const_name>‹unsigned›, T1 --> @{typ int}) $ t1)
end
| bv_term_parser (SMTLIB.Sym "sbv_to_int",[t1]) =
let
val T1 = fastype_of t1
in
SOME (Const (\<^const_name>‹signed›, T1 --> @{typ int}) $ t1)
end
| bv_term_parser (SMTLIB.Sym "bvult", [t1, t2]) =
SOME (HOLogic.mk_binrel \<^const_name>‹Orderings.less› (t1, t2))
| bv_term_parser (SMTLIB.Sym "bvule", [t1, t2]) =
SOME (HOLogic.mk_binrel \<^const_name>‹Orderings.less_eq› (t1, t2))
| bv_term_parser (SMTLIB.Sym "bvadd", t) =
let val (xs, a) = split_last t
in
SOME (fold (curry (HOLogic.mk_binop \<^const_name>‹Groups.plus›)) xs a)
end
| bv_term_parser (SMTLIB.Sym "bvsub", t) =
let val (xs, a) = split_last t
in
SOME (fold (curry (HOLogic.mk_binop \<^const_name>‹Groups.minus›)) xs a)
end
| bv_term_parser (SMTLIB.Sym "bvmul", (t::ts)) =
let
val T = fastype_of t
in
SOME (mk_lassoc (fn t1 => fn t2 => (Const ( \<^const_name>‹Groups.times›, T --> T --> T)) $ t1 $ t2) t ts)
end
| bv_term_parser (SMTLIB.Sym "bvor", (t::ts)) =
let
val T = fastype_of t
in
SOME (mk_lassoc (fn t1 => fn t2 => (Const ( \<^const_name>‹semiring_bit_operations_class.or›, T --> T --> T)) $ t1 $ t2) t ts)
end
| bv_term_parser (SMTLIB.Sym "bvxor", (t::ts)) =
let
val T = fastype_of t
in
SOME (mk_lassoc (fn t1 => fn t2 => (Const ( \<^const_name>‹semiring_bit_operations_class.xor›, T --> T --> T)) $ t1 $ t2) t ts)
end
| bv_term_parser (SMTLIB.Sym "bvxnor", [t1, t2]) =
SOME (mk_unary \<^const_name>‹ring_bit_operations_class.not› (HOLogic.mk_binop \<^const_name>‹semiring_bit_operations_class.xor› (t1, t2)))
| bv_term_parser (SMTLIB.S [SMTLIB.Sym "_", SMTLIB.Sym "zero_extend", SMTLIB.Num i], [t]) =
mk_zero_extend i t
| bv_term_parser (SMTLIB.S [SMTLIB.Sym "_", SMTLIB.Sym "sign_extend", SMTLIB.Num i], [t]) =
SOME (mk_scast i t)
| bv_term_parser (SMTLIB.Sym "sign_extend", [t1, t2]) =
SOME (Const (\<^const_name>‹Word.signed_cast›, fastype_of t2 --> dummyT) $ t2)
| bv_term_parser (SMTLIB.Sym "zero_extend", [t1, t2]) =
SOME (Const (\<^const_name>‹Word.unsigned›, fastype_of t2 --> dummyT) $ t2)
| bv_term_parser (SMTLIB.Sym "concat", t1 :: t2 :: ts) =
let
fun cat_op I1 I2 O = (Const ( \<^const_name>‹word_cat›, I1 --> I2 --> O))
fun res_type I1 I2 = (dest_wordT I1 + dest_wordT I2) |> mk_wordT
fun bitwidthConcat [x] = (x,(fastype_of x))
| bitwidthConcat (x::xs) =
let
val (curr_t, currT) = bitwidthConcat xs
val xT = fastype_of x
val resT = res_type xT currT
in
(((cat_op xT currT resT) $ x $ curr_t),resT)
end
val (res,_) = bitwidthConcat (t1::t2::ts)
in (SOME (res)) end
| bv_term_parser xs = NONE
fun bv_type_parser (SMTLIB.S [SMTLIB.Sym "_", SMTLIB.Sym "BitVec", SMTLIB.Num x], []) =
SOME (Word_Lib.mk_wordT x)
| bv_type_parser _ = NONE
val _ = Theory.setup (Context.theory_map (
SMTLIB_Interface.add_logic (30, smtlib_logic) #>
setup_builtins))
val _ = Theory.setup (Context.theory_map (
SMTLIB_Proof.add_type_parser bv_type_parser #>
SMTLIB_Proof.add_term_parser bv_term_parser))
end