packages feed

duckdb-simple-0.2.0.0: test/CancellationTests.hs

{-# LANGUAGE BlockArguments #-}
{-# LANGUAGE LambdaCase #-}
{-# LANGUAGE OverloadedStrings #-}
{-# LANGUAGE ScopedTypeVariables #-}
{-# OPTIONS_GHC -Wno-deprecations #-}

-- | Cancellation during native execution must finish before handles are released.
module CancellationTests (cancellationTests) where

import Control.Concurrent (MVar, ThreadId, forkIOWithUnmask, newEmptyMVar, putMVar, readMVar, rtsSupportsBoundThreads, threadDelay, throwTo, tryPutMVar, tryReadMVar)
import Control.Exception (AsyncException (..), SomeException, bracket, fromException, mask_, try)
import Control.Monad (unless, void)
import Data.IORef (atomicWriteIORef, newIORef, readIORef)
import Data.Int (Int64)
import Database.DuckDB.FFI (c_duckdb_interrupt, c_duckdb_query)
import Database.DuckDB.Simple
import Database.DuckDB.Simple.Arrow (foldArrow_)
import qualified Database.DuckDB.Simple.Deprecated.Streaming as Streaming
import Database.DuckDB.Simple.Internal (withConnectionHandle, withQueryCString, withResult)
import GHC.Conc (BlockReason (..), ThreadStatus (..), threadStatus)
import System.Timeout (timeout)
import Test.Tasty (TestTree, testGroup)
import Test.Tasty.HUnit

-- | Exercise each native execution entry point and two cancellation races.
cancellationTests :: TestTree
cancellationTests =
    testGroup "native cancellation" $
        [ testCase name $ withConnectionWithConfig ":memory:" [("threads", "1")] \conn -> do
            assertBool "cancellation requires the threaded runtime" rtsSupportsBoundThreads
            entered <- newEmptyMVar
            createFunction conn "query_started" (signal entered >> pure (1 :: Int64))
            withCaller conn (pure ()) (run conn longQuery) \caller done -> do
                await "native query did not start" (readMVar entered)
                (_, sent) <- startCaller (throwTo caller UserInterrupt)
                outcome <- await "native query did not cancel" (readMVar done)
                assertAsync [UserInterrupt] outcome
                await "interrupt sender did not finish" (readMVar sent) >>= assertSucceeded
            assertReusable conn
        | (name, run) <-
            [ ("query_", \conn sql -> void (query_ conn sql :: IO [Only Double]))
            , ("query with a prepared statement", \conn sql -> void (query conn sql () :: IO [Only Double]))
            , ("execute_", \conn sql -> void (execute_ conn sql))
            , ("execute with a prepared statement", \conn sql -> void (execute conn sql ()))
            , ("statement cursor", \conn sql -> withStatement conn sql \stmt -> void (nextRow stmt :: IO (Maybe (Only Double))))
            , ("fold_", \conn sql -> void (fold_ conn sql (0 :: Double) (\acc (Only value) -> pure (acc + value))))
            , ("Arrow fold", \conn sql -> foldArrow_ conn sql () (\() _ _ -> pure ()))
            ]
        ]
            <> [ testCase name $ withConnectionWithConfig ":memory:" [("threads", "1")] \conn -> do
                    void (execute_ conn "SET streaming_buffer_size = '64KB'")
                    filtering <- newIORef False
                    entered <- newEmptyMVar
                    createFunction conn "discard_remaining" \(_ :: Int64) -> do
                        active <- readIORef filtering
                        if active then signal entered >> pure True else pure False
                    let sql = "SELECT i FROM range(1000000000000) t(i) WHERE NOT discard_remaining(i)"
                    withCaller conn (pure ()) (run conn sql (atomicWriteIORef filtering True)) \caller done -> do
                        await "fetch did not start after the first delivered batch" (readMVar entered)
                        (_, sent) <- startCaller (throwTo caller UserInterrupt)
                        await "streaming fetch did not cancel" (readMVar done) >>= assertAsync [UserInterrupt]
                        await "interrupt sender did not finish" (readMVar sent) >>= assertSucceeded
                    assertReusable conn
               | (name, run) <-
                    [ ("cancel a native row fetch", \conn sql delivered -> Streaming.fold_ conn sql () (\() (Only (_ :: Int64)) -> delivered))
                    , ("cancel a native Arrow fetch", \conn sql delivered -> Streaming.foldArrow_ conn sql () (\() _ _ -> delivered))
                    ,
                        ( "mixed cursor entry points retain interruptible fetching"
                        , \conn sql delivered -> withStatement conn sql \stmt -> do
                            Streaming.nextRow stmt >>= (@?= Just (Only (0 :: Int64)))
                            delivered
                            let consume =
                                    nextRow stmt >>= \case
                                        Nothing -> pure ()
                                        Just (Only (_ :: Int64)) -> consume
                            consume
                        )
                    ]
               ]
            <> [ testCase "cancellation before native entry is not cleared by startup" $
                    withConnectionWithConfig ":memory:" [("threads", "1")] \conn -> do
                        queued <- newEmptyMVar
                        begin <- newEmptyMVar
                        createFunction conn "query_started" (pure 1 :: IO Int64)
                        let action = withConnectionHandle conn \handle ->
                                withQueryCString longQuery \sql ->
                                    withResult
                                        conn
                                        longQuery
                                        (\result -> signal queued >> readMVar begin >> c_duckdb_query handle sql result)
                                        (const (pure ()))
                        withCaller conn (signal begin) action \caller done -> do
                            await "query worker did not reach its entry gate" (readMVar queued)
                            (_, sent) <- startCaller (throwTo caller UserInterrupt)
                            await "first interrupt was not delivered" (readMVar sent) >>= assertSucceeded
                            awaitCleanup caller
                            signal begin
                            await "startup cleared the cancellation" (readMVar done) >>= assertAsync [UserInterrupt]
                        assertReusable conn
               , testCase "a second cancellation waits for a blocked callback to return" $
                    withConnectionWithConfig ":memory:" [("threads", "1")] \conn -> do
                        entered <- newEmptyMVar
                        release <- newEmptyMVar
                        createFunction conn "blocked_callback" (signal entered >> readMVar release >> pure (1 :: Int64))
                        withCaller conn (signal release) (void (query_ conn "SELECT blocked_callback()" :: IO [Only Int64])) \caller done -> do
                            await "callback did not start" (readMVar entered)
                            (_, firstSent) <- startCaller (throwTo caller UserInterrupt)
                            await "first interrupt was not delivered" (readMVar firstSent) >>= assertSucceeded
                            secondAttempted <- newEmptyMVar
                            (_, secondSent) <- startCaller (signal secondAttempted >> throwTo caller ThreadKilled)
                            await "second interrupt sender did not start" (readMVar secondAttempted)
                            premature <- timeout 50000 (readMVar secondSent)
                            case premature of
                                Nothing -> pure ()
                                Just _ -> assertFailure "second interrupt escaped cleanup"
                            tryReadMVar done >>= \case
                                Nothing -> pure ()
                                Just _ -> assertFailure "query returned while its callback was still running"
                            signal release
                            await "callback cancellation did not finish" (readMVar done) >>= assertAsync [UserInterrupt, ThreadKilled]
                            await "second interrupt sender did not finish" (readMVar secondSent) >>= assertSucceeded
                        assertReusable conn
               ]

-- | Run one volatile callback before a CPU query that cannot finish in the test window.
longQuery :: Query
longQuery =
    "WITH started AS MATERIALIZED (SELECT query_started() AS seed) \
    \SELECT sum(sin((a.i + b.j + started.seed)::DOUBLE)) \
    \FROM started, range(1000000) a(i), range(1000000) b(j)"

-- | Publish a thread's outcome while asynchronous exceptions are masked.
startCaller :: IO () -> IO (ThreadId, MVar (Either SomeException ()))
startCaller action = mask_ do
    done <- newEmptyMVar
    caller <- forkIOWithUnmask \unmask -> try (unmask action) >>= putMVar done
    pure (caller, done)

-- | Release gates and interrupt unfinished SQL before the connection can close.
withCaller :: Connection -> IO () -> IO () -> (ThreadId -> MVar (Either SomeException ()) -> IO a) -> IO a
withCaller conn release action use =
    bracket (startCaller action) cleanup (uncurry use)
  where
    cleanup (_, done) = do
        release
        await "emergency query cleanup did not finish" (stop done)
    stop done = do
        finished <- tryReadMVar done
        case finished of
            Just _ -> pure ()
            Nothing -> do
                withConnectionHandle conn c_duckdb_interrupt
                threadDelay 10000
                stop done

-- | Wait until the cancelled caller has reached its interrupt-and-join loop.
awaitCleanup :: ThreadId -> IO ()
awaitCleanup caller = awaitCondition "caller did not enter cancellation cleanup" do
    threadStatus caller >>= \case
        ThreadBlocked BlockedOnForeignCall -> pure False
        ThreadBlocked _ -> pure True
        ThreadFinished -> assertFailure "caller finished before the native worker" >> pure False
        ThreadDied -> assertFailure "caller died before publishing its result" >> pure False
        ThreadRunning -> pure False

-- | Bound waits in the test thread, independently of the cancelled caller.
await :: String -> IO a -> IO a
await message action = timeout 5000000 action >>= maybe (assertFailure message >> fail message) pure

-- | Wait for a thread state without relying on an arbitrary scheduling delay.
awaitCondition :: String -> IO Bool -> IO ()
awaitCondition message condition = await message loop
  where
    loop = do
        ready <- condition
        unless ready (threadDelay 1000 >> loop)

-- | Open a test gate once; cleanup can safely call this again.
signal :: MVar () -> IO ()
signal gate = void (tryPutMVar gate ())

-- | Require the requested asynchronous exception to reach the caller.
assertAsync :: [AsyncException] -> Either SomeException () -> Assertion
assertAsync expected = \case
    Left err -> case fromException err of
        Just actual -> assertBool ("unexpected asynchronous exception: " <> show actual) (actual `elem` expected)
        Nothing -> assertFailure ("expected asynchronous cancellation, got " <> show err)
    Right () -> assertFailure "query completed instead of reporting cancellation"

-- | Check that an interrupt sender completed normally.
assertSucceeded :: Either SomeException () -> Assertion
assertSucceeded = either (assertFailure . show) pure

-- | The interrupted native query must have released its connection state.
assertReusable :: Connection -> Assertion
assertReusable conn = do
    rows <- query_ conn "SELECT 42" :: IO [Only Int64]
    rows @?= [Only 42]