-- SPDX-FileCopyrightText: 2025 Sören Tempel <soeren+git@soeren-tempel.net>
--
-- SPDX-License-Identifier: MIT AND GPL-3.0-only
{-# LANGUAGE PatternSynonyms #-}

module SimpleBV
  ( SExpr,
    SMT.Solver,
    SMT.defaultConfig,
    SMT.newLogger,
    SMT.newLoggerWithHandle,
    SMT.newSolver,
    SMT.newSolverWithConfig,
    SMT.solverLogger,
    SMT.smtSolverLogger,
    SMT.setLogic,
    SMT.push,
    SMT.pop,
    SMT.popMany,
    SMT.check,
    SMT.Result (..),
    SMT.Value (..),
    pattern W,
    pattern Byte,
    pattern Half,
    pattern Word,
    pattern Long,
    width,
    const,
    declareBV,
    assert,
    sexprToVal,
    getValue,
    getValues,
    toSMT,
    ite,
    and,
    or,
    not,
    eq,
    bvLit,
    bvAdd,
    bvAShr,
    bvLShr,
    bvAnd,
    bvMul,
    bvNeg,
    bvOr,
    bvSDiv,
    bvSLeq,
    bvSLt,
    bvSGeq,
    bvSGt,
    bvSRem,
    bvShl,
    bvSub,
    bvUDiv,
    bvULeq,
    bvUGeq,
    bvUGt,
    bvULt,
    bvURem,
    bvXOr,
    concat,
    extract,
    signExtend,
    zeroExtend,
  )
where

import Control.DeepSeq (NFData, NFData1)
import Data.Bits (shiftL, shiftR, (.&.))
import GHC.Generics (Generic, Generic1)
import SimpleSMT qualified as SMT
import Prelude hiding (and, concat, const, not, or)

data Expr a
  = Var String
  | Int Integer
  | And a a
  | Or a a
  | Neg a
  | Not a
  | Eq a a
  | BvAdd a a
  | BvAShr a a
  | BvLShr a a
  | BvAnd a a
  | BvMul a a
  | BvOr a a
  | BvSDiv a a
  | BvSLeq a a
  | BvSLt a a
  | BvSGeq a a
  | BvSGt a a
  | BvSRem a a
  | BvShl a a
  | BvSub a a
  | BvUDiv a a
  | BvULeq a a
  | BvUGeq a a
  | BvUGt a a
  | BvULt a a
  | BvURem a a
  | BvXOr a a
  | Concat a a
  | Ite a a a
  | Extract Int Int a
  | SignExtend Integer a
  | ZeroExtend Integer a
  deriving (Int -> Expr a -> ShowS
[Expr a] -> ShowS
Expr a -> String
(Int -> Expr a -> ShowS)
-> (Expr a -> String) -> ([Expr a] -> ShowS) -> Show (Expr a)
forall a. Show a => Int -> Expr a -> ShowS
forall a. Show a => [Expr a] -> ShowS
forall a. Show a => Expr a -> String
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: forall a. Show a => Int -> Expr a -> ShowS
showsPrec :: Int -> Expr a -> ShowS
$cshow :: forall a. Show a => Expr a -> String
show :: Expr a -> String
$cshowList :: forall a. Show a => [Expr a] -> ShowS
showList :: [Expr a] -> ShowS
Show, Expr a -> Expr a -> Bool
(Expr a -> Expr a -> Bool)
-> (Expr a -> Expr a -> Bool) -> Eq (Expr a)
forall a. Eq a => Expr a -> Expr a -> Bool
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: forall a. Eq a => Expr a -> Expr a -> Bool
== :: Expr a -> Expr a -> Bool
$c/= :: forall a. Eq a => Expr a -> Expr a -> Bool
/= :: Expr a -> Expr a -> Bool
Eq, (forall x. Expr a -> Rep (Expr a) x)
-> (forall x. Rep (Expr a) x -> Expr a) -> Generic (Expr a)
forall x. Rep (Expr a) x -> Expr a
forall x. Expr a -> Rep (Expr a) x
forall a.
(forall x. a -> Rep a x) -> (forall x. Rep a x -> a) -> Generic a
forall a x. Rep (Expr a) x -> Expr a
forall a x. Expr a -> Rep (Expr a) x
$cfrom :: forall a x. Expr a -> Rep (Expr a) x
from :: forall x. Expr a -> Rep (Expr a) x
$cto :: forall a x. Rep (Expr a) x -> Expr a
to :: forall x. Rep (Expr a) x -> Expr a
Generic, (forall a. Expr a -> Rep1 Expr a)
-> (forall a. Rep1 Expr a -> Expr a) -> Generic1 Expr
forall a. Rep1 Expr a -> Expr a
forall a. Expr a -> Rep1 Expr a
forall k (f :: k -> *).
(forall (a :: k). f a -> Rep1 f a)
-> (forall (a :: k). Rep1 f a -> f a) -> Generic1 f
$cfrom1 :: forall a. Expr a -> Rep1 Expr a
from1 :: forall a. Expr a -> Rep1 Expr a
$cto1 :: forall a. Rep1 Expr a -> Expr a
to1 :: forall a. Rep1 Expr a -> Expr a
Generic1)

instance (NFData a) => NFData (Expr a)

instance NFData1 Expr

data SExpr
  = SExpr
  { SExpr -> Int
width :: Int,
    SExpr -> Expr SExpr
sexpr :: Expr SExpr
  }
  deriving (Int -> SExpr -> ShowS
[SExpr] -> ShowS
SExpr -> String
(Int -> SExpr -> ShowS)
-> (SExpr -> String) -> ([SExpr] -> ShowS) -> Show SExpr
forall a.
(Int -> a -> ShowS) -> (a -> String) -> ([a] -> ShowS) -> Show a
$cshowsPrec :: Int -> SExpr -> ShowS
showsPrec :: Int -> SExpr -> ShowS
$cshow :: SExpr -> String
show :: SExpr -> String
$cshowList :: [SExpr] -> ShowS
showList :: [SExpr] -> ShowS
Show, SExpr -> SExpr -> Bool
(SExpr -> SExpr -> Bool) -> (SExpr -> SExpr -> Bool) -> Eq SExpr
forall a. (a -> a -> Bool) -> (a -> a -> Bool) -> Eq a
$c== :: SExpr -> SExpr -> Bool
== :: SExpr -> SExpr -> Bool
$c/= :: SExpr -> SExpr -> Bool
/= :: SExpr -> SExpr -> Bool
Eq, (forall x. SExpr -> Rep SExpr x)
-> (forall x. Rep SExpr x -> SExpr) -> Generic SExpr
forall x. Rep SExpr x -> SExpr
forall x. SExpr -> Rep SExpr x
forall a.
(forall x. a -> Rep a x) -> (forall x. Rep a x -> a) -> Generic a
$cfrom :: forall x. SExpr -> Rep SExpr x
from :: forall x. SExpr -> Rep SExpr x
$cto :: forall x. Rep SExpr x -> SExpr
to :: forall x. Rep SExpr x -> SExpr
Generic)

instance NFData SExpr

toSMT :: SExpr -> SMT.SExpr
toSMT :: SExpr -> SExpr
toSMT SExpr
expr =
  case SExpr -> Expr SExpr
sexpr SExpr
expr of
    (Var String
name) -> String -> SExpr
SMT.const String
name
    (Int Integer
v) -> [SExpr] -> SExpr
SMT.List [String -> SExpr
SMT.Atom String
"_", String -> SExpr
SMT.Atom (String
"bv" String -> ShowS
forall a. [a] -> [a] -> [a]
++ Integer -> String
forall a. Show a => a -> String
show Integer
v), String -> SExpr
SMT.Atom (String -> SExpr) -> String -> SExpr
forall a b. (a -> b) -> a -> b
$ Int -> String
forall a. Show a => a -> String
show (SExpr -> Int
width SExpr
expr)]
    (Or SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.or (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (Ite SExpr
cond SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr -> SExpr
SMT.ite (SExpr -> SExpr
toSMT SExpr
cond) (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (And SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.and (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (Not SExpr
v) -> SExpr -> SExpr
SMT.not (SExpr -> SExpr
toSMT SExpr
v)
    (Neg SExpr
v) -> SExpr -> SExpr
SMT.bvNeg (SExpr -> SExpr
toSMT SExpr
v)
    (SignExtend Integer
n SExpr
v) -> Integer -> SExpr -> SExpr
SMT.signExtend Integer
n (SExpr -> SExpr
toSMT SExpr
v)
    (ZeroExtend Integer
n SExpr
v) -> Integer -> SExpr -> SExpr
SMT.zeroExtend Integer
n (SExpr -> SExpr
toSMT SExpr
v)
    (Eq SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.eq (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (Concat SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.concat (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (Extract Int
o Int
w SExpr
e) -> SExpr -> Integer -> Integer -> SExpr
SMT.extract (SExpr -> SExpr
toSMT SExpr
e) (Int -> Integer
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Integer) -> Int -> Integer
forall a b. (a -> b) -> a -> b
$ Int
o Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
w Int -> Int -> Int
forall a. Num a => a -> a -> a
- Int
1) (Int -> Integer
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
o)
    (BvAnd SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvAnd (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvAShr SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvAShr (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvLShr SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvLShr (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvAdd SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvAdd (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvMul SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvMul (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvOr SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvOr (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvSDiv SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvSDiv (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvSLeq SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvSLeq (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvSLt SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvSLt (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvSGeq SExpr
lhs SExpr
rhs) -> String -> [SExpr] -> SExpr
SMT.fun String
"bvsge" [SExpr -> SExpr
toSMT SExpr
lhs, SExpr -> SExpr
toSMT SExpr
rhs]
    (BvSGt SExpr
lhs SExpr
rhs) -> String -> [SExpr] -> SExpr
SMT.fun String
"bvsgt" [SExpr -> SExpr
toSMT SExpr
lhs, SExpr -> SExpr
toSMT SExpr
rhs]
    (BvSRem SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvSRem (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvShl SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvShl (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvSub SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvSub (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvUDiv SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvUDiv (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvULeq SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvULeq (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvUGeq SExpr
lhs SExpr
rhs) -> String -> [SExpr] -> SExpr
SMT.fun String
"bvuge" [SExpr -> SExpr
toSMT SExpr
lhs, SExpr -> SExpr
toSMT SExpr
rhs]
    (BvUGt SExpr
lhs SExpr
rhs) -> String -> [SExpr] -> SExpr
SMT.fun String
"bvugt" [SExpr -> SExpr
toSMT SExpr
lhs, SExpr -> SExpr
toSMT SExpr
rhs]
    (BvULt SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvULt (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvURem SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvURem (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)
    (BvXOr SExpr
lhs SExpr
rhs) -> SExpr -> SExpr -> SExpr
SMT.bvXOr (SExpr -> SExpr
toSMT SExpr
lhs) (SExpr -> SExpr
toSMT SExpr
rhs)

boolWidth :: Int
boolWidth :: Int
boolWidth = Int
1

pattern E :: Expr SExpr -> SExpr
pattern $mE :: forall {r}. SExpr -> (Expr SExpr -> r) -> ((# #) -> r) -> r
E expr <- SExpr {sexpr = expr, width = _}

pattern W :: Int -> SExpr
pattern $mW :: forall {r}. SExpr -> (Int -> r) -> ((# #) -> r) -> r
W w <- SExpr {width = w}

pattern Byte :: SExpr
pattern $mByte :: forall {r}. SExpr -> ((# #) -> r) -> ((# #) -> r) -> r
Byte <- SExpr {width = 8}

pattern Half :: SExpr
pattern $mHalf :: forall {r}. SExpr -> ((# #) -> r) -> ((# #) -> r) -> r
Half <- SExpr {width = 16}

pattern Word :: SExpr
pattern $mWord :: forall {r}. SExpr -> ((# #) -> r) -> ((# #) -> r) -> r
Word <- SExpr {width = 32}

pattern Long :: SExpr
pattern $mLong :: forall {r}. SExpr -> ((# #) -> r) -> ((# #) -> r) -> r
Long <- SExpr {width = 64}

------------------------------------------------------------------------

const :: String -> Int -> SExpr
const :: String -> Int -> SExpr
const String
name Int
width = Int -> Expr SExpr -> SExpr
SExpr Int
width (String -> Expr SExpr
forall a. String -> Expr a
Var String
name)

declareBV :: SMT.Solver -> String -> Int -> IO SExpr
declareBV :: Solver -> String -> Int -> IO SExpr
declareBV Solver
solver String
name Int
width = do
  let bits :: SExpr
bits = Integer -> SExpr
SMT.tBits (Integer -> SExpr) -> Integer -> SExpr
forall a b. (a -> b) -> a -> b
$ Int -> Integer
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
width
  Solver -> String -> SExpr -> IO SExpr
SMT.declare Solver
solver String
name SExpr
bits IO SExpr -> IO SExpr -> IO SExpr
forall a b. IO a -> IO b -> IO b
forall (m :: * -> *) a b. Monad m => m a -> m b -> m b
>> SExpr -> IO SExpr
forall a. a -> IO a
forall (f :: * -> *) a. Applicative f => a -> f a
pure (String -> Int -> SExpr
const String
name Int
width)

bvLit :: Int -> Integer -> SExpr
bvLit :: Int -> Integer -> SExpr
bvLit Int
width Integer
value = Int -> Expr SExpr -> SExpr
SExpr Int
width (Integer -> Expr SExpr
forall a. Integer -> Expr a
Int Integer
value)

sexprToVal :: SExpr -> SMT.Value
sexprToVal :: SExpr -> Value
sexprToVal (E (Var String
n)) = SExpr -> Value
SMT.Other (SExpr -> Value) -> SExpr -> Value
forall a b. (a -> b) -> a -> b
$ String -> SExpr
SMT.Atom String
n
sexprToVal e :: SExpr
e@(E (Int Integer
i)) = Int -> Integer -> Value
SMT.Bits (SExpr -> Int
width SExpr
e) Integer
i
sexprToVal SExpr
_ = SExpr -> Value
SMT.Other (SExpr -> Value) -> SExpr -> Value
forall a b. (a -> b) -> a -> b
$ String -> SExpr
SMT.Atom String
"_"

assert :: SMT.Solver -> SExpr -> IO ()
assert :: Solver -> SExpr -> IO ()
assert Solver
solver = Solver -> SExpr -> IO ()
SMT.assert Solver
solver (SExpr -> IO ()) -> (SExpr -> SExpr) -> SExpr -> IO ()
forall b c a. (b -> c) -> (a -> b) -> a -> c
. SExpr -> SExpr
toSMT

getValue :: SMT.Solver -> SExpr -> IO SMT.Value
getValue :: Solver -> SExpr -> IO Value
getValue Solver
solver = Solver -> SExpr -> IO Value
SMT.getExpr Solver
solver (SExpr -> IO Value) -> (SExpr -> SExpr) -> SExpr -> IO Value
forall b c a. (b -> c) -> (a -> b) -> a -> c
. SExpr -> SExpr
toSMT

getValues :: SMT.Solver -> [SExpr] -> IO [(String, SMT.Value)]
getValues :: Solver -> [SExpr] -> IO [(String, Value)]
getValues Solver
solver [SExpr]
exprs = do
  ((SExpr, Value) -> (String, Value))
-> [(SExpr, Value)] -> [(String, Value)]
forall a b. (a -> b) -> [a] -> [b]
map (SExpr, Value) -> (String, Value)
go ([(SExpr, Value)] -> [(String, Value)])
-> IO [(SExpr, Value)] -> IO [(String, Value)]
forall (f :: * -> *) a b. Functor f => (a -> b) -> f a -> f b
<$> Solver -> [SExpr] -> IO [(SExpr, Value)]
SMT.getExprs Solver
solver ((SExpr -> SExpr) -> [SExpr] -> [SExpr]
forall a b. (a -> b) -> [a] -> [b]
map SExpr -> SExpr
toSMT [SExpr]
exprs)
  where
    go :: (SMT.SExpr, SMT.Value) -> (String, SMT.Value)
    go :: (SExpr, Value) -> (String, Value)
go (SMT.Atom String
name, Value
value) = (String
name, Value
value)
    go (SExpr, Value)
_ = String -> (String, Value)
forall a. HasCallStack => String -> a
error String
"non-atomic variable in inputVars"

---------------------------------------------------------------------------

ite :: SExpr -> SExpr -> SExpr -> SExpr
ite :: SExpr -> SExpr -> SExpr -> SExpr
ite SExpr
cond SExpr
ifT SExpr
ifF = Int -> Expr SExpr -> SExpr
SExpr (SExpr -> Int
width SExpr
ifT) (SExpr -> SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> a -> Expr a
Ite SExpr
cond SExpr
ifT SExpr
ifF)

not :: SExpr -> SExpr
not :: SExpr -> SExpr
not (E (Not SExpr
cond)) = SExpr
cond
not SExpr
expr = SExpr
expr {sexpr = Not expr}

and :: SExpr -> SExpr -> SExpr
and :: SExpr -> SExpr -> SExpr
and SExpr
lhs SExpr
rhs = SExpr
lhs {sexpr = And lhs rhs}

or :: SExpr -> SExpr -> SExpr
or :: SExpr -> SExpr -> SExpr
or SExpr
lhs SExpr
rhs = SExpr
lhs {sexpr = Or lhs rhs}

signExtend :: Integer -> SExpr -> SExpr
signExtend :: Integer -> SExpr -> SExpr
signExtend Integer
n SExpr
expr = Int -> Expr SExpr -> SExpr
SExpr (SExpr -> Int
width SExpr
expr Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Integer -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral Integer
n) (Expr SExpr -> SExpr) -> Expr SExpr -> SExpr
forall a b. (a -> b) -> a -> b
$ Integer -> SExpr -> Expr SExpr
forall a. Integer -> a -> Expr a
SignExtend Integer
n SExpr
expr

zeroExtend :: Integer -> SExpr -> SExpr
zeroExtend :: Integer -> SExpr -> SExpr
zeroExtend Integer
n SExpr
expr = Int -> Expr SExpr -> SExpr
SExpr (SExpr -> Int
width SExpr
expr Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Integer -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral Integer
n) (Expr SExpr -> SExpr) -> Expr SExpr -> SExpr
forall a b. (a -> b) -> a -> b
$ Integer -> SExpr -> Expr SExpr
forall a. Integer -> a -> Expr a
ZeroExtend Integer
n SExpr
expr

------------------------------------------------------------------------

eq' :: SExpr -> SExpr -> SExpr
eq' :: SExpr -> SExpr -> SExpr
eq' SExpr
lhs SExpr
rhs = Int -> Expr SExpr -> SExpr
SExpr Int
boolWidth (Expr SExpr -> SExpr) -> Expr SExpr -> SExpr
forall a b. (a -> b) -> a -> b
$ SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
Eq SExpr
lhs SExpr
rhs

-- Eliminates ITE expressions when comparing with constants values, this is
-- useful in the QBE context to eliminate comparisons with truth values.
eq :: SExpr -> SExpr -> SExpr
eq :: SExpr -> SExpr -> SExpr
eq lexpr :: SExpr
lexpr@(E (Ite SExpr
cond (E (Int Integer
ifT)) (E (Int Integer
ifF)))) rexpr :: SExpr
rexpr@(E (Int Integer
other))
  | Integer
other Integer -> Integer -> Bool
forall a. Eq a => a -> a -> Bool
== Integer
ifT = SExpr
cond
  | Integer
other Integer -> Integer -> Bool
forall a. Eq a => a -> a -> Bool
== Integer
ifF = SExpr -> SExpr
not SExpr
cond
  | Bool
otherwise = SExpr -> SExpr -> SExpr
eq' SExpr
lexpr SExpr
rexpr
eq SExpr
lhs SExpr
rhs = SExpr -> SExpr -> SExpr
eq' SExpr
lhs SExpr
rhs

concat' :: SExpr -> SExpr -> SExpr
concat' :: SExpr -> SExpr -> SExpr
concat' SExpr
lhs SExpr
rhs =
  Int -> Expr SExpr -> SExpr
SExpr (SExpr -> Int
width SExpr
lhs Int -> Int -> Int
forall a. Num a => a -> a -> a
+ SExpr -> Int
width SExpr
rhs) (Expr SExpr -> SExpr) -> Expr SExpr -> SExpr
forall a b. (a -> b) -> a -> b
$ SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
Concat SExpr
lhs SExpr
rhs

-- Replace 0 concats with zero extension: (concat (_ bv0 8) buf6)
concatZeros :: SExpr -> SExpr -> SExpr
concatZeros :: SExpr -> SExpr -> SExpr
concatZeros lhs :: SExpr
lhs@(E (Int Integer
0)) SExpr
rhs = Integer -> SExpr -> SExpr
zeroExtend (Int -> Integer
forall a b. (Integral a, Num b) => a -> b
fromIntegral (Int -> Integer) -> Int -> Integer
forall a b. (a -> b) -> a -> b
$ SExpr -> Int
width SExpr
lhs) SExpr
rhs
concatZeros SExpr
lhs SExpr
rhs = SExpr -> SExpr -> SExpr
concat' SExpr
lhs SExpr
rhs

-- Replaces continuous concat expressions with a single extract expression.
concat :: SExpr -> SExpr -> SExpr
concat :: SExpr -> SExpr -> SExpr
concat
  lhs :: SExpr
lhs@(E (Extract Int
loff Int
lwidth latom :: SExpr
latom@(E Expr SExpr
exprLhs)))
  rhs :: SExpr
rhs@(E (Extract Int
roff Int
rwidth (E Expr SExpr
exprRhs)))
    | Expr SExpr
exprLhs Expr SExpr -> Expr SExpr -> Bool
forall a. Eq a => a -> a -> Bool
== Expr SExpr
exprRhs Bool -> Bool -> Bool
&& (Int
roff Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
rwidth) Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
loff = SExpr -> Int -> Int -> SExpr
extract SExpr
latom Int
roff (Int
lwidth Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
rwidth)
    | Bool
otherwise = SExpr -> SExpr -> SExpr
concatZeros SExpr
lhs SExpr
rhs
concat SExpr
lhs SExpr
rhs = SExpr -> SExpr -> SExpr
concatZeros SExpr
lhs SExpr
rhs

extract' :: SExpr -> Int -> Int -> SExpr
extract' :: SExpr -> Int -> Int -> SExpr
extract' SExpr
expr Int
off Int
w = Int -> Expr SExpr -> SExpr
SExpr Int
w (Expr SExpr -> SExpr) -> Expr SExpr -> SExpr
forall a b. (a -> b) -> a -> b
$ Int -> Int -> SExpr -> Expr SExpr
forall a. Int -> Int -> a -> Expr a
Extract Int
off Int
w SExpr
expr

-- Eliminate extract expression where the value already has the desired bits.
extractSameWidth :: SExpr -> Int -> Int -> SExpr
extractSameWidth :: SExpr -> Int -> Int -> SExpr
extractSameWidth SExpr
expr Int
off Int
w
  | Int
off Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0 Bool -> Bool -> Bool
&& SExpr -> Int
width SExpr
expr Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
w = SExpr
expr
  | Bool
otherwise = SExpr -> Int -> Int -> SExpr
extract' SExpr
expr Int
off Int
w

-- Eliminate nested extract expression of the same width.
extractNested :: SExpr -> Int -> Int -> SExpr
extractNested :: SExpr -> Int -> Int -> SExpr
extractNested expr :: SExpr
expr@(E (Extract Int
ioff Int
iwidth SExpr
_)) Int
off Int
width =
  if Int
ioff Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
off Bool -> Bool -> Bool
&& Int
iwidth Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
width
    then SExpr
expr
    else SExpr -> Int -> Int -> SExpr
extractSameWidth SExpr
expr Int
off Int
width
extractNested SExpr
expr Int
off Int
width = SExpr -> Int -> Int -> SExpr
extractSameWidth SExpr
expr Int
off Int
width

-- Performs direct extractions of constant immediate values.
extractConst :: SExpr -> Int -> Int -> SExpr
extractConst :: SExpr -> Int -> Int -> SExpr
extractConst (E (Int Integer
value)) Int
off Int
w =
  Int -> Expr SExpr -> SExpr
SExpr Int
w (Expr SExpr -> SExpr)
-> (Integer -> Expr SExpr) -> Integer -> SExpr
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Integer -> Expr SExpr
forall a. Integer -> Expr a
Int (Integer -> SExpr) -> Integer -> SExpr
forall a b. (a -> b) -> a -> b
$ Integer -> Int -> Integer
forall {a}. (Bits a, Num a) => a -> Int -> a
truncTo (Integer
value Integer -> Int -> Integer
forall a. Bits a => a -> Int -> a
`shiftR` Int
off) Int
w
  where
    truncTo :: a -> Int -> a
truncTo a
v Int
bits = a
v a -> a -> a
forall a. Bits a => a -> a -> a
.&. ((a
1 a -> Int -> a
forall a. Bits a => a -> Int -> a
`shiftL` Int
bits) a -> a -> a
forall a. Num a => a -> a -> a
- a
1)
extractConst SExpr
expr Int
off Int
width = SExpr -> Int -> Int -> SExpr
extractNested SExpr
expr Int
off Int
width

-- This performs constant propagation for subtyping of condition values (i.e.
-- the conversion from long to word).
extractIte :: SExpr -> Int -> Int -> SExpr
extractIte :: SExpr -> Int -> Int -> SExpr
extractIte (E (Ite SExpr
cond ifT :: SExpr
ifT@(E (Int Integer
_)) ifF :: SExpr
ifF@(E (Int Integer
_)))) Int
off Int
w =
  let ex :: SExpr -> SExpr
ex SExpr
x = SExpr -> Int -> Int -> SExpr
extractConst SExpr
x Int
off Int
w
   in Int -> Expr SExpr -> SExpr
SExpr Int
w (Expr SExpr -> SExpr) -> Expr SExpr -> SExpr
forall a b. (a -> b) -> a -> b
$ SExpr -> SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> a -> Expr a
Ite SExpr
cond (SExpr -> SExpr
ex SExpr
ifT) (SExpr -> SExpr
ex SExpr
ifF)
extractIte SExpr
expr Int
off Int
width = SExpr -> Int -> Int -> SExpr
extractConst SExpr
expr Int
off Int
width

extractZeros ::
  SExpr ->
  Int ->
  Int ->
  SExpr
extractZeros :: SExpr -> Int -> Int -> SExpr
extractZeros expr :: SExpr
expr@(E (ZeroExtend Integer
extBits SExpr
inner)) Int
exOff Int
exWidth
  | Int
exOff Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= SExpr -> Int
width SExpr
inner Bool -> Bool -> Bool
&& Integer
extBits Integer -> Integer -> Bool
forall a. Ord a => a -> a -> Bool
> Integer
0 = Int -> Integer -> SExpr
bvLit Int
exWidth Integer
0 -- only extracting zeros
  | Bool
otherwise = SExpr -> Int -> Int -> SExpr
extractIte SExpr
expr Int
exOff Int
exWidth
extractZeros SExpr
outer Int
exOff Int
exWidth = SExpr -> Int -> Int -> SExpr
extractIte SExpr
outer Int
exOff Int
exWidth

extractExt' ::
  (Integer -> SExpr -> Expr SExpr) ->
  SExpr ->
  Integer ->
  SExpr ->
  Int ->
  Int ->
  SExpr
extractExt' :: (Integer -> SExpr -> Expr SExpr)
-> SExpr -> Integer -> SExpr -> Int -> Int -> SExpr
extractExt' Integer -> SExpr -> Expr SExpr
cons SExpr
outer Integer
extBits SExpr
inner Int
exOff Int
exWidth
  -- If we are only extracting the non-extended bytes...
  | SExpr -> Int
width SExpr
inner Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
>= Int
exOff Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
exWidth = SExpr -> Int -> Int -> SExpr
extractZeros SExpr
inner Int
exOff Int
exWidth
  -- Consider: ((_ extract 31 0) ((_ zero_extend 56) byte))
  | Int
exWidth Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
< Integer -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral Integer
extBits Bool -> Bool -> Bool
&& Int
exOff Int -> Int -> Bool
forall a. Eq a => a -> a -> Bool
== Int
0 =
      Int -> Expr SExpr -> SExpr
SExpr Int
exWidth (Expr SExpr -> SExpr) -> Expr SExpr -> SExpr
forall a b. (a -> b) -> a -> b
$ Integer -> SExpr -> Expr SExpr
cons (Integer
extBits Integer -> Integer -> Integer
forall a. Num a => a -> a -> a
- Int -> Integer
forall a b. (Integral a, Num b) => a -> b
fromIntegral Int
exWidth) SExpr
inner
  -- No folding...
  | Bool
otherwise = SExpr -> Int -> Int -> SExpr
extractZeros SExpr
outer Int
exOff Int
exWidth

-- Remove ZeroExtend and SignExtend expression where we don't use
-- the extended bits because we extract below the extended size.
extractExt :: SExpr -> Int -> Int -> SExpr
extractExt :: SExpr -> Int -> Int -> SExpr
extractExt expr :: SExpr
expr@(E (SignExtend Integer
extBits SExpr
inner)) Int
exOff Int
exWidth =
  (Integer -> SExpr -> Expr SExpr)
-> SExpr -> Integer -> SExpr -> Int -> Int -> SExpr
extractExt' Integer -> SExpr -> Expr SExpr
forall a. Integer -> a -> Expr a
SignExtend SExpr
expr Integer
extBits SExpr
inner Int
exOff Int
exWidth
extractExt expr :: SExpr
expr@(E (ZeroExtend Integer
extBits SExpr
inner)) Int
exOff Int
exWidth =
  (Integer -> SExpr -> Expr SExpr)
-> SExpr -> Integer -> SExpr -> Int -> Int -> SExpr
extractExt' Integer -> SExpr -> Expr SExpr
forall a. Integer -> a -> Expr a
ZeroExtend SExpr
expr Integer
extBits SExpr
inner Int
exOff Int
exWidth
extractExt SExpr
expr Int
off Int
w = SExpr -> Int -> Int -> SExpr
extractIte SExpr
expr Int
off Int
w

extract :: SExpr -> Int -> Int -> SExpr
extract :: SExpr -> Int -> Int -> SExpr
extract = SExpr -> Int -> Int -> SExpr
extractExt

------------------------------------------------------------------------

binOp' :: (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp' :: (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp' SExpr -> SExpr -> Expr SExpr
op SExpr
lhs SExpr
rhs = SExpr
lhs {sexpr = op lhs rhs}

binOp :: (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
-- Consider: (bvslt ((_ zero_extend 24) byte0) ((_ zero_extend 24) byte1))
-- TODO: The following only works if 'op' does not consider sign-bits. Otherwise,
--       there is no semantic expression equivalence after this folding operation.
-- binOp op lhs@(E (ZeroExtend _ lhsInner)) rhs@(E (ZeroExtend _ rhsInner)) =
--   if width lhsInner == width rhsInner
--     then binOp op lhsInner rhsInner
--     else binOp' op lhs rhs
binOp :: (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
op SExpr
lhs SExpr
rhs = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp' SExpr -> SExpr -> Expr SExpr
op SExpr
lhs SExpr
rhs

-- TODO: Generate these using template-haskell.

bvNeg :: SExpr -> SExpr
bvNeg :: SExpr -> SExpr
bvNeg SExpr
x = SExpr
x {sexpr = Neg x}

bvAdd :: SExpr -> SExpr -> SExpr
bvAdd :: SExpr -> SExpr -> SExpr
bvAdd = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvAdd

bvAShr :: SExpr -> SExpr -> SExpr
bvAShr :: SExpr -> SExpr -> SExpr
bvAShr = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvAShr

bvLShr :: SExpr -> SExpr -> SExpr
bvLShr :: SExpr -> SExpr -> SExpr
bvLShr = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvLShr

bvAnd :: SExpr -> SExpr -> SExpr
bvAnd :: SExpr -> SExpr -> SExpr
bvAnd = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvAnd

bvMul :: SExpr -> SExpr -> SExpr
bvMul :: SExpr -> SExpr -> SExpr
bvMul = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvMul

bvOr :: SExpr -> SExpr -> SExpr
bvOr :: SExpr -> SExpr -> SExpr
bvOr = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvOr

bvSDiv :: SExpr -> SExpr -> SExpr
bvSDiv :: SExpr -> SExpr -> SExpr
bvSDiv = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvSDiv

bvSLeq :: SExpr -> SExpr -> SExpr
bvSLeq :: SExpr -> SExpr -> SExpr
bvSLeq = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvSLeq

bvSLt :: SExpr -> SExpr -> SExpr
bvSLt :: SExpr -> SExpr -> SExpr
bvSLt = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvSLt

bvSGeq :: SExpr -> SExpr -> SExpr
bvSGeq :: SExpr -> SExpr -> SExpr
bvSGeq = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvSGeq

bvSGt :: SExpr -> SExpr -> SExpr
bvSGt :: SExpr -> SExpr -> SExpr
bvSGt = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvSGt

bvSRem :: SExpr -> SExpr -> SExpr
bvSRem :: SExpr -> SExpr -> SExpr
bvSRem = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvSRem

bvShl :: SExpr -> SExpr -> SExpr
bvShl :: SExpr -> SExpr -> SExpr
bvShl = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvShl

bvSub :: SExpr -> SExpr -> SExpr
bvSub :: SExpr -> SExpr -> SExpr
bvSub = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvSub

bvUDiv :: SExpr -> SExpr -> SExpr
bvUDiv :: SExpr -> SExpr -> SExpr
bvUDiv = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvUDiv

bvULeq :: SExpr -> SExpr -> SExpr
bvULeq :: SExpr -> SExpr -> SExpr
bvULeq = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvULeq

bvUGeq :: SExpr -> SExpr -> SExpr
bvUGeq :: SExpr -> SExpr -> SExpr
bvUGeq = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvUGeq

bvUGt :: SExpr -> SExpr -> SExpr
bvUGt :: SExpr -> SExpr -> SExpr
bvUGt = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvUGt

bvULt :: SExpr -> SExpr -> SExpr
bvULt :: SExpr -> SExpr -> SExpr
bvULt = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvULt

bvURem :: SExpr -> SExpr -> SExpr
-- Fold constant bvURem operations which are emitted a lot in our generated
-- SMT-LIB because of QBE's "shift-value modulo bitsize"-semantics.
bvURem :: SExpr -> SExpr -> SExpr
bvURem vlhs :: SExpr
vlhs@(E (Int Integer
lhs)) (E (Int Integer
rhs)) =
  Int -> Expr SExpr -> SExpr
SExpr (SExpr -> Int
width SExpr
vlhs) (Expr SExpr -> SExpr) -> Expr SExpr -> SExpr
forall a b. (a -> b) -> a -> b
$
    -- XXX: On urem-by-zero, SMT-LIB returns the lhs.
    if Integer
rhs Integer -> Integer -> Bool
forall a. Eq a => a -> a -> Bool
== Integer
0
      then Integer -> Expr SExpr
forall a. Integer -> Expr a
Int Integer
lhs
      else Integer -> Expr SExpr
forall a. Integer -> Expr a
Int (Integer -> Expr SExpr) -> Integer -> Expr SExpr
forall a b. (a -> b) -> a -> b
$ Integer
lhs Integer -> Integer -> Integer
forall a. Integral a => a -> a -> a
`rem` Integer
rhs
bvURem SExpr
lhs SExpr
rhs = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvURem SExpr
lhs SExpr
rhs

bvXOr :: SExpr -> SExpr -> SExpr
bvXOr :: SExpr -> SExpr -> SExpr
bvXOr = (SExpr -> SExpr -> Expr SExpr) -> SExpr -> SExpr -> SExpr
binOp SExpr -> SExpr -> Expr SExpr
forall a. a -> a -> Expr a
BvXOr