packages feed

futhark-0.28.1: src/Language/Futhark/Interpreter/FFI/ServerM.hs

module Language.Futhark.Interpreter.FFI.ServerM
  ( FS.TypeName,
    ValueRef,
    Server,
    startServer,
    newServer,
    stopServer,
    ServerM,
    runServerM,
    gc,
    release,
    call,
    -- Interrogation
    inputs,
    output,
    kind,
    vtype,
    -- Primitives
    getPrim,
    putPrim,
    putData,
    getData,
    -- Arrays
    rank,
    elemType,
    mkArray,
    shape,
    index,
    -- Records
    fieldOrder,
    mkRecord,
    project,
    unzipArray,
    -- Sums
    variants,
    mkSum,
    destruct,
    -- Error handling convenience
    throwNothing,
  )
where

import Control.Exception (catch)
import Control.Monad (replicateM)
import Control.Monad.Except (ExceptT, MonadError, runExceptT, throwError)
import Control.Monad.IO.Class (MonadIO (liftIO))
import Control.Monad.Reader (ReaderT, asks, runReaderT)
import Data.IORef (IORef, atomicModifyIORef', mkWeakIORef, newIORef, readIORef)
import Data.List (intercalate)
import Data.Map qualified as M
import Data.Set qualified as S
import Data.Text qualified as T
import Data.Unique (hashUnique, newUnique)
import Data.Vector.Storable qualified as V
import Futhark.Data qualified as D
import Futhark.Server qualified as FS
import Futhark.Server.Values qualified as FS
import Futhark.Util (mapAccumLM)
import Language.Futhark.Interpreter.FFI.AtomicList as AL
import Language.Futhark.Syntax

-- | Converts a PrimValue to a Data Value
pToD :: PrimValue -> D.Value
pToD (SignedValue (Int8Value i)) = D.putValue1 i
pToD (SignedValue (Int16Value i)) = D.putValue1 i
pToD (SignedValue (Int32Value i)) = D.putValue1 i
pToD (SignedValue (Int64Value i)) = D.putValue1 i
pToD (UnsignedValue (Int8Value i)) = D.putValue1 (fromIntegral i :: Word8)
pToD (UnsignedValue (Int16Value i)) = D.putValue1 (fromIntegral i :: Word16)
pToD (UnsignedValue (Int32Value i)) = D.putValue1 (fromIntegral i :: Word32)
pToD (UnsignedValue (Int64Value i)) = D.putValue1 (fromIntegral i :: Word64)
pToD (FloatValue (Float16Value f)) = D.putValue1 f
pToD (FloatValue (Float32Value f)) = D.putValue1 f
pToD (FloatValue (Float64Value f)) = D.putValue1 f
pToD (BoolValue b) = D.putValue1 b

-- | Converts a Data Value to a PrimValue, assuming that it is a singleton
dToP :: D.Value -> PrimValue
dToP (D.I8Value _ vs) = SignedValue $ Int8Value $ vs V.! 0
dToP (D.I16Value _ vs) = SignedValue $ Int16Value $ vs V.! 0
dToP (D.I32Value _ vs) = SignedValue $ Int32Value $ vs V.! 0
dToP (D.I64Value _ vs) = SignedValue $ Int64Value $ vs V.! 0
dToP (D.U8Value _ vs) = UnsignedValue $ Int8Value $ fromIntegral $ vs V.! 0
dToP (D.U16Value _ vs) = UnsignedValue $ Int16Value $ fromIntegral $ vs V.! 0
dToP (D.U32Value _ vs) = UnsignedValue $ Int32Value $ fromIntegral $ vs V.! 0
dToP (D.U64Value _ vs) = UnsignedValue $ Int64Value $ fromIntegral $ vs V.! 0
dToP (D.F16Value _ vs) = FloatValue $ Float16Value $ vs V.! 0
dToP (D.F32Value _ vs) = FloatValue $ Float32Value $ vs V.! 0
dToP (D.F64Value _ vs) = FloatValue $ Float64Value $ vs V.! 0
dToP (D.BoolValue _ vs) = BoolValue $ vs V.! 0

newtype ValueRef = ValueRef (IORef FS.VarName)

data Server = Server
  { server :: FS.Server,
    -- | Variables whose 'ValueRef' has been garbage collected.
    queue :: AL.AtomicList FS.VarName,
    -- | Variables created by us that have not yet been freed.
    live :: IORef (S.Set FS.VarName)
  }

newtype ServerM a = ServerM (ReaderT Server (ExceptT String IO) a)
  deriving
    ( Functor,
      Applicative,
      Monad,
      MonadError String,
      MonadIO
    )

askServer :: ServerM FS.Server
askServer = ServerM $ asks server

askQueue :: ServerM (AL.AtomicList FS.VarName)
askQueue = ServerM $ asks queue

modifyLive :: (S.Set FS.VarName -> S.Set FS.VarName) -> ServerM ()
modifyLive f = do
  r <- ServerM $ asks live
  liftIO $ atomicModifyIORef' r $ (,()) . f

startServer :: FS.ServerCfg -> IO Server
startServer cfg = newServer =<< FS.startServer cfg

-- | Use an already-running server. Shutting it down remains the
-- responsibility of whoever started it.
newServer :: FS.Server -> IO Server
newServer s = Server s <$> AL.new <*> newIORef mempty

-- | Shut down the server. Returns a message on termination failure.
stopServer :: Server -> IO (Maybe T.Text)
stopServer s =
  (Nothing <$ FS.stopServer (server s))
    `catch` \(FS.ServerException e) -> pure $ Just e

runServerM :: Server -> ServerM a -> IO (Either String a)
runServerM s (ServerM m) = runExceptT $ runReaderT m s

varName :: ValueRef -> ServerM FS.VarName
varName (ValueRef r) = liftIO $ readIORef r

uniqueName :: ServerM FS.VarName
uniqueName = ("v" <>) . T.show . hashUnique <$> liftIO newUnique

mkValueRef :: FS.VarName -> ServerM ValueRef
mkValueRef n = do
  modifyLive $ S.insert n
  r <- liftIO $ newIORef n
  q <- askQueue
  _ <- liftIO $ mkWeakIORef r $ AL.prepend n q
  pure $ ValueRef r

gc :: ServerM ()
gc = freeVars =<< liftIO . AL.flush =<< askQueue

freeVars :: [FS.VarName] -> ServerM ()
freeVars vns = do
  s <- askServer
  liftIO (FS.cmdFree s vns)
    >>= throwServerJust ("cmdFree failed on variables " ++ csList (map T.unpack vns) ++ ".")
  modifyLive (`S.difference` S.fromList vns)

-- | End the use of this 'Server'. The variables of the given values are
-- adopted by the caller under the given names, and every other variable we
-- have created is freed, whether or not its 'ValueRef' is still reachable.
-- Neither the 'Server' nor any 'ValueRef' may be used afterwards. A variable
-- may occur more than once, in which case it is adopted under the first of its
-- names. Returns the name of the variable of each value.
release :: [(ValueRef, FS.VarName)] -> ServerM [FS.VarName]
release adopted = do
  s <- askServer
  -- Everything in the queue is also live, so it is freed below.
  _ <- askQueue >>= liftIO . AL.flush
  srcs <- mapM (varName . fst) adopted
  let adopt renamed (src, dst)
        | Just dst' <- M.lookup src renamed = pure (renamed, dst')
        | otherwise = do
            liftIO (FS.cmdRename s src dst)
              >>= throwServerJust ("cmdRename failed on variable " ++ T.unpack src ++ ".")
            pure (M.insert src dst renamed, dst)
  (renamed, dsts) <- mapAccumLM adopt mempty $ zip srcs $ map snd adopted
  modifyLive (`S.difference` M.keysSet renamed)
  freeVars . S.toList =<< liftIO . readIORef =<< ServerM (asks live)
  pure dsts

call :: Name -> [ValueRef] -> ServerM ValueRef
call fn ps = do
  s <- askServer
  nps <- mapM varName ps
  ndst <- uniqueName
  -- A failing call is usually the program itself failing (e.g. OOB), so report
  -- just what the server said.
  _ <-
    liftIO (FS.cmdCall s (nameToText fn) ndst nps)
      >>= either (throwError . T.unpack . T.unlines . FS.failureMsg) pure
  mkValueRef ndst

-- Interrogation
inputs :: Name -> ServerM [FS.TypeName]
inputs fn = do
  s <- askServer
  map FS.inputType <$> (liftIO (FS.cmdInputs s $ nameToText fn) >>= throwServerLeft ("cmdInputs failed on function " ++ nameToString fn ++ "."))

output :: Name -> ServerM FS.TypeName
output fn = do
  s <- askServer
  FS.outputType <$> (liftIO (FS.cmdOutput s $ nameToText fn) >>= throwServerLeft ("cmdOutput failed on function " ++ nameToString fn ++ "."))

kind :: FS.TypeName -> ServerM FS.Kind
kind tn = do
  s <- askServer
  liftIO (FS.cmdKind s tn) >>= throwServerLeft ("cmdKind failed on type " ++ T.unpack tn ++ ".")

vtype :: ValueRef -> ServerM FS.TypeName
vtype vr = do
  s <- askServer
  vn <- varName vr
  liftIO (FS.cmdType s vn) >>= throwServerLeft ("cmdType failed on variable " ++ T.unpack vn ++ ".")

-- Primitives
getPrim :: ValueRef -> ServerM PrimValue
getPrim vr = do
  s <- askServer
  nsrc <- varName vr
  v <- liftIO (FS.getValue s nsrc) >>= throwLeft ("Failed to get primitive variable " ++ T.unpack nsrc ++ ".")
  pure $ dToP v

putPrim :: PrimValue -> ServerM ValueRef
putPrim = putData . pToD

-- | Put an entire value on the server at once. This is only possible for
-- values that can be represented in the Futhark data format (primitives and
-- arrays of primitive).
putData :: D.Value -> ServerM ValueRef
putData v = do
  s <- askServer
  ndst <- uniqueName
  liftIO (FS.putValue s ndst v)
    >>= throwServerJust ("Failed to put value of type " ++ T.unpack (D.valueTypeText (D.valueType v)) ++ ".")
  mkValueRef ndst

-- Arrays
rank :: FS.TypeName -> ServerM Int
rank tn = do
  s <- askServer
  liftIO (FS.cmdRank s tn) >>= throwServerLeft ("cmdRank failed on type " ++ T.unpack tn ++ ".")

elemType :: FS.TypeName -> ServerM FS.TypeName
elemType tn = do
  s <- askServer
  liftIO (FS.cmdElemtype s tn) >>= throwServerLeft ("cmdElemtype failed on type " ++ T.unpack tn ++ ".")

mkArray :: FS.TypeName -> [Int64] -> [ValueRef] -> ServerM ValueRef
mkArray tn dims vs = do
  s <- askServer
  vns <- mapM varName vs
  dst <- uniqueName
  liftIO (FS.cmdNewArray s dst tn (map fromIntegral dims) vns) >>= throwServerJust ("cmdNewArray failed on type " ++ T.unpack tn ++ " with variables " ++ csList (map T.unpack vns) ++ ".")
  mkValueRef dst

shape :: ValueRef -> ServerM [Int64]
shape vr = do
  s <- askServer
  vn <- varName vr
  map fromIntegral <$> (liftIO (FS.cmdShape s vn) >>= throwServerLeft ("cmdShape failed on variable " ++ T.unpack vn ++ "."))

-- | Retrieve an entire value from the server at once. This is only possible for
-- values that can be represented in the Futhark data format (primitives and
-- arrays of primitive).
getData :: ValueRef -> ServerM (Maybe D.Value)
getData vr = do
  s <- askServer
  n <- varName vr
  either (const Nothing) Just <$> liftIO (FS.getValue s n)

index :: [Int64] -> ValueRef -> ServerM ValueRef
index is src = do
  s <- askServer
  nsrc <- varName src
  ndst <- uniqueName
  liftIO (FS.cmdIndex s ndst nsrc $ map fromIntegral is) >>= throwServerJust ("cmdIndex failed on source " ++ T.unpack nsrc ++ ", destination " ++ T.unpack ndst ++ ", and index " ++ show is ++ ".")
  mkValueRef ndst

-- Records

-- | The fields of a record type, in the order the server uses.
fieldOrder :: FS.TypeName -> ServerM [(Name, FS.TypeName)]
fieldOrder tn = do
  s <- askServer
  fs <- liftIO (FS.cmdFields s tn) >>= throwServerLeft ("cmdFields failed on type " ++ T.unpack tn ++ ".")
  pure $ map (\f -> (nameFromText $ FS.fieldName f, FS.fieldType f)) fs

-- | Split an array of records into one array per field, in 'fieldOrder'. The
-- fields of an array cannot be projected one element at a time, and doing so
-- would anyway be impossible for an empty array.
unzipArray :: ValueRef -> Int -> ServerM [ValueRef]
unzipArray src n = do
  s <- askServer
  nsrc <- varName src
  ndsts <- replicateM n uniqueName
  liftIO (FS.cmdUnzip s nsrc ndsts)
    >>= throwServerJust ("cmdUnzip failed on variable " ++ T.unpack nsrc ++ ".")
  mapM mkValueRef ndsts

mkRecord :: FS.TypeName -> M.Map Name ValueRef -> ServerM ValueRef
mkRecord tn vrm = do
  s <- askServer
  fns <- map (nameFromText . FS.fieldName) <$> (liftIO (FS.cmdFields s tn) >>= throwServerLeft ("cmdFields failed on type " ++ T.unpack tn ++ "."))
  vns <-
    mapM
      ( \fn ->
          throwNothing ("Missing field " ++ nameToString fn ++ " when constructing record of type " ++ T.unpack tn ++ ".") (M.lookup fn vrm)
            >>= varName
      )
      fns
  dst <- uniqueName
  liftIO (FS.cmdNew s dst tn vns) >>= throwServerJust ("cmdNew failed on type " ++ T.unpack tn ++ " with variables " ++ csList (map T.unpack vns) ++ ".")
  mkValueRef dst

project :: ValueRef -> Name -> ServerM ValueRef
project src fn = do
  s <- askServer
  nsrc <- varName src
  ndst <- uniqueName
  liftIO (FS.cmdProject s ndst nsrc $ nameToText fn)
    >>= throwServerJust ("cmdProject failed on source " ++ T.unpack nsrc ++ ", destination " ++ T.unpack ndst ++ ", and field " ++ nameToString fn ++ ".")
  mkValueRef ndst

-- Sums
variants :: FS.TypeName -> ServerM (M.Map Name [FS.TypeName])
variants tn = do
  s <- askServer
  vs <- liftIO (FS.cmdVariants s tn) >>= throwServerLeft ("cmdVariants failed on type " ++ T.unpack tn ++ ".")
  pure $ M.fromList $ map (\v -> (nameFromText $ FS.variantName v, FS.variantTypes v)) vs

mkSum :: FS.TypeName -> Name -> [ValueRef] -> ServerM ValueRef
mkSum tn vn vrs = do
  s <- askServer
  vns <- mapM varName vrs
  dst <- uniqueName
  liftIO (FS.cmdConstruct s dst tn (nameToText vn) vns)
    >>= throwServerJust ("cmdConstruct failed on type " ++ T.unpack tn ++ ", variant " ++ nameToString vn ++ " with variables " ++ csList (map T.unpack vns) ++ ".")
  mkValueRef dst

-- | The variant of a sum, and its payload.
destruct :: ValueRef -> ServerM (Name, [ValueRef])
destruct src = do
  vn <- variant src
  tn <- vtype src
  vts <- variants tn >>= throwNothing ("Variant " ++ nameToString vn ++ " is not part of its own sum type, " ++ T.unpack tn ++ ". This should be impossible.") . M.lookup vn
  do
    s <- askServer
    nsrc <- varName src
    ndsts <- mapM (const uniqueName) vts
    liftIO (FS.cmdDestruct s nsrc ndsts)
      >>= throwServerJust ("cmdVariants failed on source " ++ T.unpack nsrc ++ ", destinations " ++ csList (map T.unpack ndsts) ++ ".")
    (vn,) <$> mapM mkValueRef ndsts

variant :: ValueRef -> ServerM Name
variant src = do
  s <- askServer
  nsrc <- varName src
  vn <-
    liftIO (FS.cmdVariant s nsrc)
      >>= throwServerLeft ("cmdIndex failed on variable " ++ T.unpack nsrc ++ ".")
  pure $ nameFromText vn

-- Error handling convenience
formatServerError :: String -> FS.CmdFailure -> String
formatServerError e f | e == mempty = formatServerError "Server error." f
formatServerError e f = T.unpack $ T.unlines $ T.pack e : "Failure message:" : FS.failureMsg f

throwServerLeft :: (MonadError String m) => String -> Either FS.CmdFailure a -> m a
throwServerLeft e (Left c) = throwError $ formatServerError e c
throwServerLeft _ (Right v) = pure v

throwServerJust :: (MonadError String m) => String -> Maybe FS.CmdFailure -> m ()
throwServerJust e c = throwJust $ formatServerError e <$> c

throwLeft :: (MonadError String m) => String -> Either T.Text a -> m a
throwLeft t (Left e) = throwError $ T.unpack $ T.unlines [T.pack t, e]
throwLeft _ (Right v) = pure v

throwJust :: (MonadError String m) => Maybe String -> m ()
throwJust (Just e) = throwError e
throwJust Nothing = pure ()

throwNothing :: (MonadError String m) => String -> Maybe a -> m a
throwNothing _ (Just v) = pure v
throwNothing e Nothing = throwError e

csList :: [String] -> String
csList = intercalate ","