packages feed

hs-bindgen-1.0.0.0: test/th/test-th.hs

{-# LANGUAGE DataKinds #-}
{-# LANGUAGE DerivingVia #-}
{-# LANGUAGE OverloadedRecordDot #-}

-- TODO <https://github.com/well-typed/hs-bindgen/issues/1575>
-- Once we do FLAMs differently no orphans should be necessary anymore.
{-# OPTIONS_GHC -Wno-orphans #-}

module Main (main) where

import Control.Exception
import Data.Proxy
import Data.Vector.Storable qualified as VS
import Foreign
import Foreign.C
import Test.Tasty
import Test.Tasty.HUnit

import HsBindgen.Runtime.ConstantArray qualified as CA
import HsBindgen.Runtime.FLAM qualified as FLAM
import HsBindgen.Runtime.IsArray qualified as IsA
import HsBindgen.Runtime.Prelude
import HsBindgen.Runtime.PtrConst qualified as PtrConst
import HsBindgen.Runtime.Union qualified as Union

import Test.Common.Util.Tasty
import Test.TH.StaticCounterA qualified as CounterA
import Test.TH.StaticCounterB qualified as CounterB
import Test.TH.Test01 qualified as Test01
import Test.TH.Test02 qualified as Test02

{-------------------------------------------------------------------------------
  Test01
-------------------------------------------------------------------------------}

-- StructBasic is generated
t01Val :: Test01.StructBasic
t01Val = Test01.StructBasic
    { Test01.structBasic_field1 = 0
    , Test01.structBasic_field2 = 1
    }

-- .. and can be poked
_t01PokeVal :: Ptr Test01.StructBasic -> IO ()
_t01PokeVal ptr = poke ptr t01Val

-- Macros
_t01MyPlus :: CLong -> CLong -> CLong
_t01MyPlus x y = Test01.pLUS x y

-- Flexible array member (orphan instance)
--
-- TODO <https://github.com/well-typed/hs-bindgen/issues/1575
-- Once we do FLAMs differently this instance should be generated by bindgen.
instance FLAM.NumElems CLong Test01.StructFLAM_Aux where
    numElems x = fromIntegral (Test01.structFLAM_length x)

-- Bounded and Enum orphan instances
--
-- TODO <https://github.com/well-typed/hs-bindgen/issues/1527>
-- This should be generated by bindgen instead.
deriving via AsCEnum Test01.EnumBasic  instance Bounded Test01.EnumBasic
deriving via AsCEnum Test01.EnumBasic  instance Enum    Test01.EnumBasic
deriving via AsCEnum Test01.EnumNeg    instance Bounded Test01.EnumNeg
deriving via AsCEnum Test01.EnumNeg    instance Enum    Test01.EnumNeg
deriving via AsCEnum Test01.EnumNonSeq instance Bounded Test01.EnumNonSeq
deriving via AsCEnum Test01.EnumNonSeq instance Enum    Test01.EnumNonSeq
deriving via AsCEnum Test01.EnumSame   instance Bounded Test01.EnumSame
deriving via AsCEnum Test01.EnumSame   instance Enum    Test01.EnumSame

-- Unit tests
test01 :: TestTree
test01 = testGroup "test_01"
    [ testCase "constants" $ do
        sizeOf (undefined :: Test01.StructBasic) @?= 8
        alignment (undefined :: Test01.StructBasic) @?= 4

        sizeOf (undefined :: Test01.StructFixedSizeArray) @?= 12
        alignment (undefined :: Test01.StructFixedSizeArray)  @?= 4

    , testCase "ConstantArray peek-poke-roundtrip" $ do
        let s = Test01.StructFixedSizeArray 5 (CA.repeat 12)
        s' <- alloca $ \ptr -> do
            poke ptr s
            peek ptr

        s' @?= s

    , testCase "Bitfield" $ do
        let s = Test01.StructBitfield 5 1 1 2
        s' <- alloca $ \ptr -> do
            poke ptr s
            peek ptr
        s' @?= s

    , testCase "function" $ do
        res <- Test01.my_fma 2 3 5
        res @?= 11

    , testCase "flam" $ do
        let n = 10
        bracket (Test01.flam_alloc n) Test01.flam_free $ \ptr -> do
            -- Peek.
            struct <- readRaw ptr
            struct.flam @?= VS.fromList [0..9]

            -- Poke, Ok.
            let v' = VS.fromList [10..19]
            writeRaw ptr (struct { FLAM.flam = v' })
            struct' <- readRaw ptr
            struct'.flam @?= v'

            -- Poke, error.
            let vLengthMismatch = VS.fromList [0]
            assertException "Expected FlamLengthMismatch" (Proxy :: Proxy FLAM.FlamLengthMismatch) $
              writeRaw ptr (struct { FLAM.flam = vLengthMismatch })
            struct'' <- readRaw ptr
            struct''.flam @?= v'

            -- Reverse.
            structBeforeReverse <- readRaw ptr
            Test01.reverse ptr
            structAfterReverse <- readRaw ptr
            assertEqual "reversed flam"
              (structBeforeReverse.flam)
              (VS.reverse structAfterReverse.flam)

    , testCase "EnumBasic" $
        [minBound..maxBound]
          @?= [Test01.ENUM_BASIC_A, Test01.ENUM_BASIC_B, Test01.ENUM_BASIC_C]

    , testCase "EnumNeg" $
        [minBound..maxBound]
          @?= [Test01.ENUM_NEG_A, Test01.ENUM_NEG_B, Test01.ENUM_NEG_C]

    , testCase "EnumNonSeq" $
        [minBound..maxBound]
          @?= [ Test01.ENUM_NON_SEQ_A
              , Test01.ENUM_NON_SEQ_B
              , Test01.ENUM_NON_SEQ_C
              ]

    , testCase "EnumSame" $ do
        Test01.ENUM_SAME_B @?= Test01.ENUM_SAME_C
        [minBound..maxBound]
          @?= [Test01.ENUM_SAME_A, Test01.ENUM_SAME_B, Test01.ENUM_SAME_D]

    , testCase "struct-arg" $ do
        res1 <- thing_fun_1 (Test01.Thing 10)
        10 @?= res1

        res2 <- Test01.thing_fun_1 (Test01.Thing 20)
        20 @?= res2

    , testCase "struct-res" $ do
        res <- thing_fun_2 11
        Test01.Thing 11 @?= res

        res' <- Test01.thing_fun_2 22
        Test01.Thing 22 @?= res'

    , testCase "struct-arg-res" $ do
        res <- thing_fun_3 (Test01.Thing 6)
        Test01.Thing 12 @?= res

        res' <- Test01.thing_fun_3 (Test01.Thing 6)
        Test01.Thing 12 @?= res'

    , testCase "union" $ do
        let x :: Test01.LongLongOrDouble
            x = Union.zero
            l :: CLLong
            l = x.longLongOrDouble_l
            d :: CDouble
            d = x.longLongOrDouble_d
        l @?= 0
        d @?= 0
        let d' :: CDouble
            d' = (Union.set @"longLongOrDouble_d" @Test01.LongLongOrDouble 3.5).longLongOrDouble_d
            l' :: CLLong
            l' = (Union.set @"longLongOrDouble_l" @Test01.LongLongOrDouble 10).longLongOrDouble_l
        l' @?= 10
        d' @?= 3.5

    , testCase "fixed-size-array" $ do
        let v = CA.repeat 4 :: ConstantArray 3 CInt
        res <- IsA.withElemPtr v (\ptr -> Test01.sum3 5 (PtrConst.unsafeFromPtr ptr))
        21 @?= res
        [4,4,4] @?= CA.toList v -- modification in sum3 aren't visible in original array.

        res' <- IsA.withElemPtr v (\ptr -> Test01.sum3 7 (PtrConst.unsafeFromPtr ptr))
        23 @?= res'

        let t = Test01.Triple v
        res'' <- IsA.withElemPtr t (\ptr -> Test01.sum3b 9 (PtrConst.unsafeFromPtr ptr))
        29 @?= res''
    ]

thing_fun_1 :: Test01.Thing -> IO CInt
thing_fun_1 x = Test01.thing_fun_1 x

thing_fun_2 :: CInt -> IO Test01.Thing
thing_fun_2 x = Test01.thing_fun_2 x

thing_fun_3 :: Test01.Thing -> IO Test01.Thing
thing_fun_3 x = Test01.thing_fun_3 x

{-------------------------------------------------------------------------------
  Test02
-------------------------------------------------------------------------------}

-- Event is generated
t02Val :: Test02.Event
t02Val = Test02.Event
    { Test02.event_id   = 42
    , Test02.event_name = nullPtr
    , Test02.event_time = CTime 1770276088
    }

-- Unit tests
test02 :: TestTree
test02 = testGroup "test_02"
    [ testCase "Event peek-poke-roundtrip" $ do
        x <- alloca $ \ptr -> do
            poke ptr t02Val
            peek ptr
        x @?= t02Val
    ]

{-------------------------------------------------------------------------------
  Static counter
-------------------------------------------------------------------------------}

testStaticCounter :: TestTree
testStaticCounter = testGroup "static_counter"
    [ testCase "independent static state" $ do
        -- Reset both counters
        CounterA.reset
        CounterB.reset

        -- Increment counter A a few times
        a0 <- CounterA.count
        a1 <- CounterA.count
        a2 <- CounterA.count

        a0 @?= 0
        a1 @?= 1
        a2 @?= 2

        -- Counter B should be independent (still at 0)
        b0 <- CounterB.count
        b0 @?= 0

        -- And counter B increments independently
        b1 <- CounterB.count
        b1 @?= 1

        -- Counter A continues from where it left off
        a3 <- CounterA.count
        a3 @?= 3
    ]

{-------------------------------------------------------------------------------
  Main
-------------------------------------------------------------------------------}

main :: IO ()
main = defaultMain $ testGroup "test-th" [
      test01
    , test02
    , testStaticCounter
    ]