{-# LANGUAGE GADTs      #-}
{-# LANGUAGE RankNTypes #-}
{-# LANGUAGE Safe       #-}

-- | Convert modular transition systems ('TransSys') into Kind2 file
-- specifications.
module Copilot.Theorem.Kind2.Translate
  ( toKind2
  ) where

import Copilot.Theorem.TransSys
import qualified Copilot.Theorem.Kind2.AST as K

import Data.Function (on)
import Data.Maybe (fromJust)

import Data.List (partition, sort, sortBy)
import Data.Map ((!))

import qualified Data.Map as Map
import qualified Data.Bimap as Bimap

-- | Produce a Kind2 file that checks the properties specified.
toKind2 :: [PropId]  -- ^ Assumptions
        -> [PropId]  -- ^ Properties to be checked
        -> TransSys  -- ^ Modular transition system holding the system spec
        -> K.File
toKind2 :: [NodeId] -> [NodeId] -> TransSys -> File
toKind2 [NodeId]
assumptions [NodeId]
checkedProps TransSys
spec =
  TransSys -> [NodeId] -> File -> File
addAssumptions TransSys
spec' [NodeId]
assumptions (File -> File) -> File -> File
forall a b. (a -> b) -> a -> b
$ TransSys -> [NodeId] -> File
trSpec TransSys
spec' [NodeId]
checkedProps
  where
    spec' :: TransSys
spec' = TransSys -> TransSys
inline TransSys
spec

trSpec :: TransSys -> [PropId] -> K.File
trSpec :: TransSys -> [NodeId] -> File
trSpec TransSys
spec [NodeId]
checkedProps = [Node] -> Node -> [Prop] -> File
K.File [Node]
otherNodes Node
topNode [Prop]
props
  where
    (Node
topNode, [Node]
otherNodes) =
      case (Node -> Bool) -> [Node] -> ([Node], [Node])
forall a. (a -> Bool) -> [a] -> ([a], [a])
partition ((NodeId -> NodeId -> Bool
forall a. Eq a => a -> a -> Bool
== TransSys -> NodeId
specTopNodeId TransSys
spec) (NodeId -> Bool) -> (Node -> NodeId) -> Node -> Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Node -> NodeId
K.nodeId) [Node]
nodes of
        ([Node
top], [Node]
others) -> (Node
top, [Node]
others)
        ([Node], [Node])
_ -> NodeId -> (Node, [Node])
forall a. HasCallStack => NodeId -> a
error (NodeId -> (Node, [Node])) -> NodeId -> (Node, [Node])
forall a b. (a -> b) -> a -> b
$ NodeId
"Kind2.Translate: top node "
                       NodeId -> NodeId -> NodeId
forall a. [a] -> [a] -> [a]
++ TransSys -> NodeId
specTopNodeId TransSys
spec NodeId -> NodeId -> NodeId
forall a. [a] -> [a] -> [a]
++ NodeId
" not found"
    nodes :: [Node]
nodes = (Node -> Node) -> [Node] -> [Node]
forall a b. (a -> b) -> [a] -> [b]
map (TransSys -> Node -> Node
trNode TransSys
spec) (TransSys -> [Node]
specNodes TransSys
spec)
    props :: [Prop]
props = ((NodeId, ExtVar) -> Prop) -> [(NodeId, ExtVar)] -> [Prop]
forall a b. (a -> b) -> [a] -> [b]
map (NodeId, ExtVar) -> Prop
trProp ([(NodeId, ExtVar)] -> [Prop]) -> [(NodeId, ExtVar)] -> [Prop]
forall a b. (a -> b) -> a -> b
$
      ((NodeId, ExtVar) -> Bool)
-> [(NodeId, ExtVar)] -> [(NodeId, ExtVar)]
forall a. (a -> Bool) -> [a] -> [a]
filter ((NodeId -> [NodeId] -> Bool
forall a. Eq a => a -> [a] -> Bool
forall (t :: * -> *) a. (Foldable t, Eq a) => a -> t a -> Bool
`elem` [NodeId]
checkedProps) (NodeId -> Bool)
-> ((NodeId, ExtVar) -> NodeId) -> (NodeId, ExtVar) -> Bool
forall b c a. (b -> c) -> (a -> b) -> a -> c
. (NodeId, ExtVar) -> NodeId
forall a b. (a, b) -> a
fst) ([(NodeId, ExtVar)] -> [(NodeId, ExtVar)])
-> [(NodeId, ExtVar)] -> [(NodeId, ExtVar)]
forall a b. (a -> b) -> a -> b
$
        Map NodeId ExtVar -> [(NodeId, ExtVar)]
forall k a. Map k a -> [(k, a)]
Map.toList (Map NodeId ExtVar -> [(NodeId, ExtVar)])
-> Map NodeId ExtVar -> [(NodeId, ExtVar)]
forall a b. (a -> b) -> a -> b
$ ((ExtVar, Prop) -> ExtVar)
-> Map NodeId (ExtVar, Prop) -> Map NodeId ExtVar
forall a b k. (a -> b) -> Map k a -> Map k b
Map.map (ExtVar, Prop) -> ExtVar
forall a b. (a, b) -> a
fst (Map NodeId (ExtVar, Prop) -> Map NodeId ExtVar)
-> Map NodeId (ExtVar, Prop) -> Map NodeId ExtVar
forall a b. (a -> b) -> a -> b
$ TransSys -> Map NodeId (ExtVar, Prop)
specProps TransSys
spec

trProp :: (PropId, ExtVar) -> K.Prop
trProp :: (NodeId, ExtVar) -> Prop
trProp (NodeId
pId, ExtVar
var) = NodeId -> Term -> Prop
K.Prop NodeId
pId (Var -> Term
trVar (Var -> Term) -> (ExtVar -> Var) -> ExtVar -> Term
forall b c a. (b -> c) -> (a -> b) -> a -> c
. ExtVar -> Var
extVarLocalPart (ExtVar -> Term) -> ExtVar -> Term
forall a b. (a -> b) -> a -> b
$ ExtVar
var)

trNode :: TransSys -> Node -> K.Node
trNode :: TransSys -> Node -> Node
trNode TransSys
spec Node
node = K.Node
  { nodeId :: NodeId
K.nodeId        = Node -> NodeId
nodeId Node
node
  , nodeStateVars :: [StateVarDef]
K.nodeStateVars = TransSys -> Node -> [StateVarDef]
gatherPredStateVars TransSys
spec Node
node
  , nodeInit :: Term
K.nodeInit      = [Term] -> Term
mkConj ([Term] -> Term) -> [Term] -> Term
forall a b. (a -> b) -> a -> b
$ Node -> [Term]
initLocals  Node
node
                               [Term] -> [Term] -> [Term]
forall a. [a] -> [a] -> [a]
++ (Expr Bool -> Term) -> [Expr Bool] -> [Term]
forall a b. (a -> b) -> [a] -> [b]
map (Bool -> Expr Bool -> Term
forall t. Bool -> Expr t -> Term
trExpr Bool
False) (Node -> [Expr Bool]
nodeConstrs Node
node)
  , nodeTrans :: Term
K.nodeTrans     = [Term] -> Term
mkConj ([Term] -> Term) -> [Term] -> Term
forall a b. (a -> b) -> a -> b
$ Node -> [Term]
transLocals Node
node
                               [Term] -> [Term] -> [Term]
forall a. [a] -> [a] -> [a]
++ (Expr Bool -> Term) -> [Expr Bool] -> [Term]
forall a b. (a -> b) -> [a] -> [b]
map (Bool -> Expr Bool -> Term
forall t. Bool -> Expr t -> Term
trExpr Bool
True) (Node -> [Expr Bool]
nodeConstrs Node
node)
  }

-- | Add the assumptions to the top node of the file by conjoining the
-- corresponding property variables to its initial state and transition
-- relation predicates.
addAssumptions :: TransSys -> [PropId] -> K.File -> K.File
addAssumptions :: TransSys -> [NodeId] -> File -> File
addAssumptions TransSys
spec [NodeId]
assumptions File
file =
  File
file { K.fileTopNode = aux (K.fileTopNode file) }
  where
    aux :: Node -> Node
aux Node
node =
      let init' :: Term
init'  = [Term] -> Term
mkConj ( Node -> Term
K.nodeInit  Node
node Term -> [Term] -> [Term]
forall a. a -> [a] -> [a]
: (NodeId -> Term) -> [NodeId] -> [Term]
forall a b. (a -> b) -> [a] -> [b]
map NodeId -> Term
K.StateVar [NodeId]
vars )
          trans' :: Term
trans' = [Term] -> Term
mkConj ( Node -> Term
K.nodeTrans Node
node Term -> [Term] -> [Term]
forall a. a -> [a] -> [a]
: (NodeId -> Term) -> [NodeId] -> [Term]
forall a b. (a -> b) -> [a] -> [b]
map NodeId -> Term
K.PrimedStateVar [NodeId]
vars )
      in Node
node { K.nodeInit = init', K.nodeTrans = trans' }

    toExtVar :: NodeId -> ExtVar
toExtVar NodeId
a = (ExtVar, Prop) -> ExtVar
forall a b. (a, b) -> a
fst ((ExtVar, Prop) -> ExtVar) -> (ExtVar, Prop) -> ExtVar
forall a b. (a -> b) -> a -> b
$ Maybe (ExtVar, Prop) -> (ExtVar, Prop)
forall a. HasCallStack => Maybe a -> a
fromJust (Maybe (ExtVar, Prop) -> (ExtVar, Prop))
-> Maybe (ExtVar, Prop) -> (ExtVar, Prop)
forall a b. (a -> b) -> a -> b
$ NodeId -> Map NodeId (ExtVar, Prop) -> Maybe (ExtVar, Prop)
forall k a. Ord k => k -> Map k a -> Maybe a
Map.lookup NodeId
a (Map NodeId (ExtVar, Prop) -> Maybe (ExtVar, Prop))
-> Map NodeId (ExtVar, Prop) -> Maybe (ExtVar, Prop)
forall a b. (a -> b) -> a -> b
$ TransSys -> Map NodeId (ExtVar, Prop)
specProps TransSys
spec
    vars :: [NodeId]
vars = (NodeId -> NodeId) -> [NodeId] -> [NodeId]
forall a b. (a -> b) -> [a] -> [b]
map (Var -> NodeId
varName (Var -> NodeId) -> (NodeId -> Var) -> NodeId -> NodeId
forall b c a. (b -> c) -> (a -> b) -> a -> c
. ExtVar -> Var
extVarLocalPart (ExtVar -> Var) -> (NodeId -> ExtVar) -> NodeId -> Var
forall b c a. (b -> c) -> (a -> b) -> a -> c
. NodeId -> ExtVar
toExtVar) [NodeId]
assumptions

-- The ordering really matters here because the variables
-- have to be given in this order in a pred call
-- Our convention :
-- * First the local variables, sorted by alphabetical order
-- * Then the imported variables, by alphabetical order on
--   the father node then by alphabetical order on the variable name

gatherPredStateVars :: TransSys -> Node -> [K.StateVarDef]
gatherPredStateVars :: TransSys -> Node -> [StateVarDef]
gatherPredStateVars TransSys
spec Node
node = [StateVarDef]
locals [StateVarDef] -> [StateVarDef] -> [StateVarDef]
forall a. [a] -> [a] -> [a]
++ [StateVarDef]
imported
  where
    nodesMap :: Map NodeId Node
nodesMap = [(NodeId, Node)] -> Map NodeId Node
forall k a. Ord k => [(k, a)] -> Map k a
Map.fromList [(Node -> NodeId
nodeId Node
n, Node
n) | Node
n <- TransSys -> [Node]
specNodes TransSys
spec]
    extVarType :: ExtVar -> K.Type
    extVarType :: ExtVar -> Type
extVarType (ExtVar NodeId
n Var
v) =
      case Node -> Map Var VarDescr
nodeLocalVars (Map NodeId Node
nodesMap Map NodeId Node -> NodeId -> Node
forall k a. Ord k => Map k a -> k -> a
! NodeId
n) Map Var VarDescr -> Var -> VarDescr
forall k a. Ord k => Map k a -> k -> a
! Var
v of
        VarDescr Type t
Integer VarDef t
_ -> Type
K.Int
        VarDescr Type t
Bool    VarDef t
_ -> Type
K.Bool
        VarDescr Type t
Real    VarDef t
_ -> Type
K.Real

    locals :: [StateVarDef]
locals =
      (Var -> StateVarDef) -> [Var] -> [StateVarDef]
forall a b. (a -> b) -> [a] -> [b]
map (\Var
v -> NodeId -> Type -> [StateVarFlag] -> StateVarDef
K.StateVarDef (Var -> NodeId
varName Var
v)
              (ExtVar -> Type
extVarType (ExtVar -> Type) -> ExtVar -> Type
forall a b. (a -> b) -> a -> b
$ NodeId -> Var -> ExtVar
ExtVar (Node -> NodeId
nodeId Node
node) Var
v) [])
         ([Var] -> [StateVarDef])
-> (Map Var VarDescr -> [Var]) -> Map Var VarDescr -> [StateVarDef]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. [Var] -> [Var]
forall a. Ord a => [a] -> [a]
sort ([Var] -> [Var])
-> (Map Var VarDescr -> [Var]) -> Map Var VarDescr -> [Var]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Map Var VarDescr -> [Var]
forall k a. Map k a -> [k]
Map.keys (Map Var VarDescr -> [StateVarDef])
-> Map Var VarDescr -> [StateVarDef]
forall a b. (a -> b) -> a -> b
$ Node -> Map Var VarDescr
nodeLocalVars Node
node

    imported :: [StateVarDef]
imported =
      ((Var, ExtVar) -> StateVarDef) -> [(Var, ExtVar)] -> [StateVarDef]
forall a b. (a -> b) -> [a] -> [b]
map (\(Var
v, ExtVar
ev) -> NodeId -> Type -> [StateVarFlag] -> StateVarDef
K.StateVarDef (Var -> NodeId
varName Var
v) (ExtVar -> Type
extVarType ExtVar
ev) [])
      ([(Var, ExtVar)] -> [StateVarDef])
-> (Bimap Var ExtVar -> [(Var, ExtVar)])
-> Bimap Var ExtVar
-> [StateVarDef]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. ((Var, ExtVar) -> (Var, ExtVar) -> Ordering)
-> [(Var, ExtVar)] -> [(Var, ExtVar)]
forall a. (a -> a -> Ordering) -> [a] -> [a]
sortBy (ExtVar -> ExtVar -> Ordering
forall a. Ord a => a -> a -> Ordering
compare (ExtVar -> ExtVar -> Ordering)
-> ((Var, ExtVar) -> ExtVar)
-> (Var, ExtVar)
-> (Var, ExtVar)
-> Ordering
forall b c a. (b -> b -> c) -> (a -> b) -> a -> a -> c
`on` (Var, ExtVar) -> ExtVar
forall a b. (a, b) -> b
snd) ([(Var, ExtVar)] -> [(Var, ExtVar)])
-> (Bimap Var ExtVar -> [(Var, ExtVar)])
-> Bimap Var ExtVar
-> [(Var, ExtVar)]
forall b c a. (b -> c) -> (a -> b) -> a -> c
. Bimap Var ExtVar -> [(Var, ExtVar)]
forall a b. Bimap a b -> [(a, b)]
Bimap.toList (Bimap Var ExtVar -> [StateVarDef])
-> Bimap Var ExtVar -> [StateVarDef]
forall a b. (a -> b) -> a -> b
$ Node -> Bimap Var ExtVar
nodeImportedVars Node
node

mkConj :: [K.Term] -> K.Term
mkConj :: [Term] -> Term
mkConj []  = Type Bool -> Bool -> Term
forall t. Type t -> t -> Term
trConst Type Bool
Bool Bool
True
mkConj [Term
x] = Term
x
mkConj [Term]
xs  = NodeId -> [Term] -> Term
K.FunApp NodeId
"and" [Term]
xs

mkEquality :: K.Term -> K.Term -> K.Term
mkEquality :: Term -> Term -> Term
mkEquality Term
t1 Term
t2 = NodeId -> [Term] -> Term
K.FunApp NodeId
"=" [Term
t1, Term
t2]

trVar :: Var -> K.Term
trVar :: Var -> Term
trVar Var
v = NodeId -> Term
K.StateVar (Var -> NodeId
varName Var
v)

trPrimedVar :: Var -> K.Term
trPrimedVar :: Var -> Term
trPrimedVar Var
v = NodeId -> Term
K.PrimedStateVar (Var -> NodeId
varName Var
v)

trConst :: Type t -> t -> K.Term
trConst :: forall t. Type t -> t -> Term
trConst Type t
Integer t
v     = NodeId -> Term
K.ValueLiteral (t -> NodeId
forall a. Show a => a -> NodeId
show t
v)
trConst Type t
Real    t
v     = NodeId -> Term
K.ValueLiteral (t -> NodeId
forall a. Show a => a -> NodeId
show t
v)
trConst Type t
Bool    t
Bool
True  = NodeId -> Term
K.ValueLiteral NodeId
"true"
trConst Type t
Bool    t
Bool
False = NodeId -> Term
K.ValueLiteral NodeId
"false"

initLocals :: Node -> [K.Term]
initLocals :: Node -> [Term]
initLocals Node
node =
  ((Var, VarDescr) -> [Term]) -> [(Var, VarDescr)] -> [Term]
forall (t :: * -> *) a b. Foldable t => (a -> [b]) -> t a -> [b]
concatMap (Var, VarDescr) -> [Term]
f (Map Var VarDescr -> [(Var, VarDescr)]
forall k a. Map k a -> [(k, a)]
Map.toList (Map Var VarDescr -> [(Var, VarDescr)])
-> Map Var VarDescr -> [(Var, VarDescr)]
forall a b. (a -> b) -> a -> b
$ Node -> Map Var VarDescr
nodeLocalVars Node
node)
  where
    f :: (Var, VarDescr) -> [Term]
f (Var
v, VarDescr Type t
t VarDef t
def) =
      case VarDef t
def of
        Pre     t
c Var
_ -> [Term -> Term -> Term
mkEquality (Var -> Term
trVar Var
v) (Type t -> t -> Term
forall t. Type t -> t -> Term
trConst Type t
t t
c)]
        Expr    Expr t
e   -> [Term -> Term -> Term
mkEquality (Var -> Term
trVar Var
v) (Bool -> Expr t -> Term
forall t. Bool -> Expr t -> Term
trExpr Bool
False Expr t
e)]
        Constrs [Expr Bool]
cs  -> (Expr Bool -> Term) -> [Expr Bool] -> [Term]
forall a b. (a -> b) -> [a] -> [b]
map (Bool -> Expr Bool -> Term
forall t. Bool -> Expr t -> Term
trExpr Bool
False) [Expr Bool]
cs

transLocals :: Node -> [K.Term]
transLocals :: Node -> [Term]
transLocals Node
node =
  ((Var, VarDescr) -> [Term]) -> [(Var, VarDescr)] -> [Term]
forall (t :: * -> *) a b. Foldable t => (a -> [b]) -> t a -> [b]
concatMap (Var, VarDescr) -> [Term]
f (Map Var VarDescr -> [(Var, VarDescr)]
forall k a. Map k a -> [(k, a)]
Map.toList (Map Var VarDescr -> [(Var, VarDescr)])
-> Map Var VarDescr -> [(Var, VarDescr)]
forall a b. (a -> b) -> a -> b
$ Node -> Map Var VarDescr
nodeLocalVars Node
node)
  where
   f :: (Var, VarDescr) -> [Term]
f (Var
v, VarDescr Type t
_ VarDef t
def) =
      case VarDef t
def of
        Pre t
_ Var
v' -> [Term -> Term -> Term
mkEquality (Var -> Term
trPrimedVar Var
v) (Var -> Term
trVar Var
v')]
        Expr Expr t
e   -> [Term -> Term -> Term
mkEquality (Var -> Term
trPrimedVar Var
v) (Bool -> Expr t -> Term
forall t. Bool -> Expr t -> Term
trExpr Bool
True Expr t
e)]
        Constrs [Expr Bool]
cs  -> (Expr Bool -> Term) -> [Expr Bool] -> [Term]
forall a b. (a -> b) -> [a] -> [b]
map (Bool -> Expr Bool -> Term
forall t. Bool -> Expr t -> Term
trExpr Bool
True) [Expr Bool]
cs

trExpr :: Bool -> Expr t -> K.Term
trExpr :: forall t. Bool -> Expr t -> Term
trExpr Bool
primed = Expr t -> Term
forall t. Expr t -> Term
tr
  where
    tr :: forall t . Expr t -> K.Term
    tr :: forall t. Expr t -> Term
tr (Const Type t
t t
c) = Type t -> t -> Term
forall t. Type t -> t -> Term
trConst Type t
t t
c
    tr (Ite Type t
_ Expr Bool
c Expr t
e1 Expr t
e2) = NodeId -> [Term] -> Term
K.FunApp NodeId
"ite" [Expr Bool -> Term
forall t. Expr t -> Term
tr Expr Bool
c, Expr t -> Term
forall t. Expr t -> Term
tr Expr t
e1, Expr t -> Term
forall t. Expr t -> Term
tr Expr t
e2]
    tr (Op1 Type t
_ Op1 t
op Expr t
e) = NodeId -> [Term] -> Term
K.FunApp (Op1 t -> NodeId
forall a. Show a => a -> NodeId
show Op1 t
op) [Expr t -> Term
forall t. Expr t -> Term
tr Expr t
e]
    tr (Op2 Type t
_ Op2 a t
op Expr a
e1 Expr a
e2) = NodeId -> [Term] -> Term
K.FunApp (Op2 a t -> NodeId
forall a. Show a => a -> NodeId
show Op2 a t
op) [Expr a -> Term
forall t. Expr t -> Term
tr Expr a
e1, Expr a -> Term
forall t. Expr t -> Term
tr Expr a
e2]
    tr (VarE Type t
_ Var
v) = if Bool
primed then Var -> Term
trPrimedVar Var
v else Var -> Term
trVar Var
v