sockets 0.1.0.0 → 0.3.0.0
raw patch · 17 files changed
+2719/−711 lines, 17 filesdep +entropydep +stmdep +textdep ~posix-apiPVP ok
version bump matches the API change (PVP)
Dependencies added: entropy, stm, text
Dependency ranges changed: posix-api
API changes (from Hackage documentation)
- Socket.Datagram.IPv4.Undestined: Accept :: Context
- Socket.Datagram.IPv4.Undestined: Bind :: Context
- Socket.Datagram.IPv4.Undestined: Close :: Context
- Socket.Datagram.IPv4.Undestined: Connect :: Context
- Socket.Datagram.IPv4.Undestined: GetName :: Context
- Socket.Datagram.IPv4.Undestined: Listen :: Context
- Socket.Datagram.IPv4.Undestined: MessageTruncated :: !Int -> !Int -> Reason
- Socket.Datagram.IPv4.Undestined: Open :: Context
- Socket.Datagram.IPv4.Undestined: Option :: Context
- Socket.Datagram.IPv4.Undestined: Receive :: Context
- Socket.Datagram.IPv4.Undestined: Send :: Context
- Socket.Datagram.IPv4.Undestined: Shutdown :: Context
- Socket.Datagram.IPv4.Undestined: SocketException :: Context -> Reason -> SocketException
- Socket.Datagram.IPv4.Undestined: [$sel:context:SocketException] :: SocketException -> Context
- Socket.Datagram.IPv4.Undestined: [$sel:reason:SocketException] :: SocketException -> Reason
- Socket.Datagram.IPv4.Undestined: data Context
- Socket.Datagram.IPv4.Undestined: data Reason
- Socket.Datagram.IPv4.Undestined: instance GHC.Classes.Eq Socket.Datagram.IPv4.Undestined.Socket
- Socket.Datagram.IPv4.Undestined: instance GHC.Classes.Ord Socket.Datagram.IPv4.Undestined.Socket
- Socket.Datagram.IPv4.Undestined: receive :: Socket -> Int -> IO (Either SocketException (Endpoint, ByteArray))
- Socket.Stream.IPv4: Accept :: Context
- Socket.Stream.IPv4: Bind :: Context
- Socket.Stream.IPv4: Close :: Context
- Socket.Stream.IPv4: Connect :: Context
- Socket.Stream.IPv4: ErrorCode :: !CInt -> Reason
- Socket.Stream.IPv4: GetName :: Context
- Socket.Stream.IPv4: Listen :: Context
- Socket.Stream.IPv4: MessageTruncated :: !Int -> !Int -> Reason
- Socket.Stream.IPv4: NegativeBytesRequested :: Reason
- Socket.Stream.IPv4: Open :: Context
- Socket.Stream.IPv4: Option :: Context
- Socket.Stream.IPv4: OptionValueSize :: Reason
- Socket.Stream.IPv4: Receive :: Context
- Socket.Stream.IPv4: RemoteNotShutdown :: Reason
- Socket.Stream.IPv4: RemoteShutdown :: Reason
- Socket.Stream.IPv4: Send :: Context
- Socket.Stream.IPv4: Shutdown :: Context
- Socket.Stream.IPv4: SocketAddressFamily :: Reason
- Socket.Stream.IPv4: SocketAddressSize :: Reason
- Socket.Stream.IPv4: SocketException :: Context -> Reason -> SocketException
- Socket.Stream.IPv4: [$sel:context:SocketException] :: SocketException -> Context
- Socket.Stream.IPv4: [$sel:reason:SocketException] :: SocketException -> Reason
- Socket.Stream.IPv4: data Context
- Socket.Stream.IPv4: data Reason
+ Socket.Datagram.IPv4.Spoof: Endpoint :: !IPv4 -> !Word16 -> Endpoint
+ Socket.Datagram.IPv4.Spoof: Message :: {-# UNPACK #-} !Endpoint -> !ByteArray -> Message
+ Socket.Datagram.IPv4.Spoof: Socket :: Fd -> Socket
+ Socket.Datagram.IPv4.Spoof: [$sel:address:Endpoint] :: Endpoint -> !IPv4
+ Socket.Datagram.IPv4.Spoof: [$sel:port:Endpoint] :: Endpoint -> !Word16
+ Socket.Datagram.IPv4.Spoof: [SendBroadcasted] :: SendException i
+ Socket.Datagram.IPv4.Spoof: [SendInterrupted] :: SendException 'Interruptible
+ Socket.Datagram.IPv4.Spoof: [SendTruncated] :: !Int -> SendException i
+ Socket.Datagram.IPv4.Spoof: [SocketFileDescriptorLimit] :: SocketException
+ Socket.Datagram.IPv4.Spoof: [SocketPermissionDenied] :: SocketException
+ Socket.Datagram.IPv4.Spoof: [payload] :: Message -> !ByteArray
+ Socket.Datagram.IPv4.Spoof: [remote] :: Message -> {-# UNPACK #-} !Endpoint
+ Socket.Datagram.IPv4.Spoof: data Endpoint
+ Socket.Datagram.IPv4.Spoof: data Message
+ Socket.Datagram.IPv4.Spoof: data SendException :: Interruptibility -> Type
+ Socket.Datagram.IPv4.Spoof: data SocketException :: Type
+ Socket.Datagram.IPv4.Spoof: instance GHC.Classes.Eq Socket.Datagram.IPv4.Spoof.Socket
+ Socket.Datagram.IPv4.Spoof: instance GHC.Classes.Ord Socket.Datagram.IPv4.Spoof.Socket
+ Socket.Datagram.IPv4.Spoof: instance GHC.Exception.Type.Exception Socket.Datagram.IPv4.Spoof.SocketException
+ Socket.Datagram.IPv4.Spoof: instance GHC.Show.Show Socket.Datagram.IPv4.Spoof.Socket
+ Socket.Datagram.IPv4.Spoof: instance GHC.Show.Show Socket.Datagram.IPv4.Spoof.SocketException
+ Socket.Datagram.IPv4.Spoof: newtype Socket
+ Socket.Datagram.IPv4.Spoof: sendMutableByteArray :: Socket -> Endpoint -> Endpoint -> MutableByteArray RealWorld -> Int -> Int -> IO (Either (SendException 'Uninterruptible) ())
+ Socket.Datagram.IPv4.Spoof: withSocket :: (Socket -> IO a) -> IO (Either SocketException a)
+ Socket.Datagram.IPv4.Undestined: Message :: {-# UNPACK #-} !Endpoint -> !ByteArray -> Message
+ Socket.Datagram.IPv4.Undestined: ReceivedMessageTruncated :: !Int -> SocketException
+ Socket.Datagram.IPv4.Undestined: ReceptionAbandoned :: SocketException
+ Socket.Datagram.IPv4.Undestined: SentMessageTruncated :: !Int -> SocketException
+ Socket.Datagram.IPv4.Undestined: [payload] :: Message -> !ByteArray
+ Socket.Datagram.IPv4.Undestined: [remote] :: Message -> {-# UNPACK #-} !Endpoint
+ Socket.Datagram.IPv4.Undestined: data Message
+ Socket.Datagram.IPv4.Undestined: receiveByteArray :: Socket -> Int -> IO (Either (ReceiveException 'Uninterruptible) Message)
+ Socket.Datagram.IPv4.Undestined: receiveMany :: Socket -> Int -> Int -> IO (Either SocketException (Array Message))
+ Socket.Datagram.IPv4.Undestined: receiveManyUnless :: STM () -> Socket -> Int -> Int -> IO (Either SocketException (Array Message))
+ Socket.Datagram.IPv4.Undestined: sendMutableByteArraySlice :: Socket -> Endpoint -> MutableByteArray RealWorld -> Int -> Int -> IO (Either (SendException 'Uninterruptible) ())
+ Socket.Stream.IPv4: Interruptible :: Interruptibility
+ Socket.Stream.IPv4: Uninterruptible :: Interruptibility
+ Socket.Stream.IPv4: [AcceptConnectionAborted] :: AcceptException i
+ Socket.Stream.IPv4: [AcceptFileDescriptorLimit] :: AcceptException i
+ Socket.Stream.IPv4: [AcceptFirewalled] :: AcceptException i
+ Socket.Stream.IPv4: [AcceptInterrupted] :: AcceptException 'Interruptible
+ Socket.Stream.IPv4: [ClosePeerContinuedSending] :: CloseException
+ Socket.Stream.IPv4: [ConnectEphemeralPortsExhausted] :: ConnectException i
+ Socket.Stream.IPv4: [ConnectFileDescriptorLimit] :: ConnectException i
+ Socket.Stream.IPv4: [ConnectFirewalled] :: ConnectException i
+ Socket.Stream.IPv4: [ConnectInterrupted] :: ConnectException 'Interruptible
+ Socket.Stream.IPv4: [ConnectNetworkUnreachable] :: ConnectException i
+ Socket.Stream.IPv4: [ConnectRefused] :: ConnectException i
+ Socket.Stream.IPv4: [ConnectTimeout] :: ConnectException i
+ Socket.Stream.IPv4: [ReceiveInterrupted] :: ReceiveException 'Interruptible
+ Socket.Stream.IPv4: [ReceiveReset] :: ReceiveException i
+ Socket.Stream.IPv4: [ReceiveShutdown] :: ReceiveException i
+ Socket.Stream.IPv4: [SendInterrupted] :: SendException 'Interruptible
+ Socket.Stream.IPv4: [SendReset] :: SendException i
+ Socket.Stream.IPv4: [SendShutdown] :: SendException i
+ Socket.Stream.IPv4: [SocketAddressInUse] :: SocketException
+ Socket.Stream.IPv4: [SocketEphemeralPortsExhausted] :: SocketException
+ Socket.Stream.IPv4: [SocketFileDescriptorLimit] :: SocketException
+ Socket.Stream.IPv4: [SocketPermissionDenied] :: SocketException
+ Socket.Stream.IPv4: accept :: Listener -> IO (Either (AcceptException 'Uninterruptible) (Connection, Endpoint))
+ Socket.Stream.IPv4: connect :: Endpoint -> IO (Either (ConnectException 'Uninterruptible) Connection)
+ Socket.Stream.IPv4: data AcceptException :: Interruptibility -> Type
+ Socket.Stream.IPv4: data CloseException :: Type
+ Socket.Stream.IPv4: data ConnectException :: Interruptibility -> Type
+ Socket.Stream.IPv4: data Interruptibility
+ Socket.Stream.IPv4: data ReceiveException :: Interruptibility -> Type
+ Socket.Stream.IPv4: data SendException :: Interruptibility -> Type
+ Socket.Stream.IPv4: disconnect :: Connection -> IO (Either CloseException ())
+ Socket.Stream.IPv4: disconnect_ :: Connection -> IO ()
+ Socket.Stream.IPv4: interruptibleAccept :: TVar Bool -> Listener -> IO (Either (AcceptException 'Interruptible) (Connection, Endpoint))
+ Socket.Stream.IPv4: interruptibleForkAcceptedUnmasked :: TVar Int -> TVar Bool -> Listener -> (Either CloseException () -> a -> IO ()) -> (Connection -> Endpoint -> IO a) -> IO (Either (AcceptException 'Interruptible) ThreadId)
+ Socket.Stream.IPv4: interruptibleReceiveBoundedMutableByteArraySlice :: TVar Bool -> Connection -> Int -> MutableByteArray RealWorld -> Int -> IO (Either (ReceiveException 'Interruptible) Int)
+ Socket.Stream.IPv4: interruptibleReceiveByteArray :: TVar Bool -> Connection -> Int -> IO (Either (ReceiveException 'Interruptible) ByteArray)
+ Socket.Stream.IPv4: interruptibleSendByteArray :: TVar Bool -> Connection -> ByteArray -> IO (Either (SendException 'Interruptible) ())
+ Socket.Stream.IPv4: interruptibleSendByteArraySlice :: TVar Bool -> Connection -> ByteArray -> Int -> Int -> IO (Either (SendException 'Interruptible) ())
+ Socket.Stream.IPv4: interruptibleSendMutableByteArraySlice :: TVar Bool -> Connection -> MutableByteArray RealWorld -> Int -> Int -> IO (Either (SendException 'Interruptible) ())
+ Socket.Stream.IPv4: listen :: Endpoint -> IO (Either SocketException (Listener, Word16))
+ Socket.Stream.IPv4: receiveBoundedMutableByteArraySlice :: Connection -> Int -> MutableByteArray RealWorld -> Int -> IO (Either (ReceiveException 'Uninterruptible) Int)
+ Socket.Stream.IPv4: unlisten :: Listener -> IO ()
+ Socket.Stream.IPv4: unlisten_ :: Listener -> IO ()
- Socket.Datagram.IPv4.Undestined: ErrorCode :: !CInt -> Reason
+ Socket.Datagram.IPv4.Undestined: ErrorCode :: !CInt -> SocketException
- Socket.Datagram.IPv4.Undestined: NegativeBytesRequested :: Reason
+ Socket.Datagram.IPv4.Undestined: NegativeBytesRequested :: SocketException
- Socket.Datagram.IPv4.Undestined: OptionValueSize :: Reason
+ Socket.Datagram.IPv4.Undestined: OptionValueSize :: SocketException
- Socket.Datagram.IPv4.Undestined: RemoteNotShutdown :: Reason
+ Socket.Datagram.IPv4.Undestined: RemoteNotShutdown :: SocketException
- Socket.Datagram.IPv4.Undestined: RemoteShutdown :: Reason
+ Socket.Datagram.IPv4.Undestined: RemoteShutdown :: SocketException
- Socket.Datagram.IPv4.Undestined: SocketAddressFamily :: Reason
+ Socket.Datagram.IPv4.Undestined: SocketAddressFamily :: !CInt -> SocketException
- Socket.Datagram.IPv4.Undestined: SocketAddressSize :: Reason
+ Socket.Datagram.IPv4.Undestined: SocketAddressSize :: SocketException
- Socket.Stream.IPv4: data SocketException
+ Socket.Stream.IPv4: data SocketException :: Type
- Socket.Stream.IPv4: forkAccepted :: Listener -> (Either SocketException a -> IO ()) -> (Connection -> Endpoint -> IO a) -> IO (Either SocketException ThreadId)
+ Socket.Stream.IPv4: forkAccepted :: Listener -> (Either CloseException () -> a -> IO ()) -> (Connection -> Endpoint -> IO a) -> IO (Either (AcceptException 'Uninterruptible) ThreadId)
- Socket.Stream.IPv4: forkAcceptedUnmasked :: Listener -> (Either SocketException a -> IO ()) -> (Connection -> Endpoint -> IO a) -> IO (Either SocketException ThreadId)
+ Socket.Stream.IPv4: forkAcceptedUnmasked :: Listener -> (Either CloseException () -> a -> IO ()) -> (Connection -> Endpoint -> IO a) -> IO (Either (AcceptException 'Uninterruptible) ThreadId)
- Socket.Stream.IPv4: receiveBoundedByteArray :: Connection -> Int -> IO (Either SocketException ByteArray)
+ Socket.Stream.IPv4: receiveBoundedByteArray :: Connection -> Int -> IO (Either (ReceiveException 'Uninterruptible) ByteArray)
- Socket.Stream.IPv4: receiveByteArray :: Connection -> Int -> IO (Either SocketException ByteArray)
+ Socket.Stream.IPv4: receiveByteArray :: Connection -> Int -> IO (Either (ReceiveException 'Uninterruptible) ByteArray)
- Socket.Stream.IPv4: receiveMutableByteArray :: Connection -> MutableByteArray RealWorld -> IO (Either SocketException ())
+ Socket.Stream.IPv4: receiveMutableByteArray :: Connection -> MutableByteArray RealWorld -> IO (Either (ReceiveException 'Uninterruptible) ())
- Socket.Stream.IPv4: sendByteArray :: Connection -> ByteArray -> IO (Either SocketException ())
+ Socket.Stream.IPv4: sendByteArray :: Connection -> ByteArray -> IO (Either (SendException 'Uninterruptible) ())
- Socket.Stream.IPv4: sendByteArraySlice :: Connection -> ByteArray -> Int -> Int -> IO (Either SocketException ())
+ Socket.Stream.IPv4: sendByteArraySlice :: Connection -> ByteArray -> Int -> Int -> IO (Either (SendException 'Uninterruptible) ())
- Socket.Stream.IPv4: sendMutableByteArray :: Connection -> MutableByteArray RealWorld -> IO (Either SocketException ())
+ Socket.Stream.IPv4: sendMutableByteArray :: Connection -> MutableByteArray RealWorld -> IO (Either (SendException 'Uninterruptible) ())
- Socket.Stream.IPv4: sendMutableByteArraySlice :: Connection -> MutableByteArray RealWorld -> Int -> Int -> IO (Either SocketException ())
+ Socket.Stream.IPv4: sendMutableByteArraySlice :: Connection -> MutableByteArray RealWorld -> Int -> Int -> IO (Either (SendException 'Uninterruptible) ())
- Socket.Stream.IPv4: withAccepted :: Listener -> (Connection -> Endpoint -> IO a) -> IO (Either SocketException a)
+ Socket.Stream.IPv4: withAccepted :: Listener -> (Either CloseException () -> a -> IO b) -> (Connection -> Endpoint -> IO a) -> IO (Either (AcceptException 'Uninterruptible) b)
- Socket.Stream.IPv4: withConnection :: Endpoint -> (Connection -> IO a) -> IO (Either SocketException a)
+ Socket.Stream.IPv4: withConnection :: Endpoint -> (Either CloseException () -> a -> IO b) -> (Connection -> IO a) -> IO (Either (ConnectException 'Uninterruptible) b)
Files
- CHANGELOG.md +5/−1
- bench/Macro.hs +131/−103
- example/Main.hs +1/−1
- sockets.cabal +32/−5
- src-debug/Socket/Debug.hs +4/−0
- src-mmsg/Socket/Datagram/IPv4/Undestined/Multiple.hs +101/−0
- src-no-mmsg/Socket/Datagram/IPv4/Undestined/Multiple.hs +140/−0
- src-production/Socket/Debug.hs +4/−0
- src/Socket.hs +100/−31
- src/Socket/Datagram.hs +44/−0
- src/Socket/Datagram/IPv4/Spoof.hs +290/−0
- src/Socket/Datagram/IPv4/Undestined.hs +179/−66
- src/Socket/Datagram/IPv4/Undestined/Internal.hs +24/−0
- src/Socket/IPv4.hs +52/−0
- src/Socket/Stream.hs +167/−0
- src/Socket/Stream/IPv4.hs +1211/−493
- test/Main.hs +234/−11
CHANGELOG.md view
@@ -1,5 +1,9 @@ # Revision history for stream-sockets -## 0.1.0.0 -- YYYY-mm-dd+## Unreleased++* introduce `receiveMutableByteArraySlice` which allow receiving in an arbitrary buffer slice.++## 0.1.0.0 -- 2019-01-18 * First version. Released on an unsuspecting world.
bench/Macro.hs view
@@ -1,144 +1,172 @@ {-# language BangPatterns #-}+{-# language DuplicateRecordFields #-}+{-# language LambdaCase #-}+{-# language MagicHash #-}+{-# language OverloadedStrings #-} {-# language ScopedTypeVariables #-}+{-# language TypeFamilies #-}+{-# language UnboxedTuples #-} -import Control.Concurrent.Async (concurrently)+-- This is a benchmark designed to stress both the sockets library+-- and the GHC event manager. It opens a moderate number of datagram sockets+-- that each belong to one of two teams: A and B. There is one worker+-- thread for each socket. The thread sends a datagram to a pseudorandomly+-- detemined socket on the other team. Then, it waits to receive a datagram+-- from a socket on the other team. All worker threads repeatedly perform+-- this task forever. Once a large number N of total receives have occurred,+-- the lucky worker thread performing the Nth receives fills an MVar that+-- tells the main thread that enough work has been done. The main thread+-- prints the total number of elapsed nanoseconds and then exits. This+-- benchmark does not attempt to close the sockets before exiting.+--+-- Here are some interesting results of running this benchmark on a+-- Dell Precision 7510 (Intel Xeon CPU E3-1505M, 4 physical cores,+-- 8 virtual cores with hyperthreading, 32GB memory):+--+-- * Wall-clock time:+-- * Non-threaded runtime: 37s+-- * Threaded runtime:+-- * N1: 12s+-- * N2: 7.5s+-- * N4: 19s+-- * Memory:+-- * Total Allocations: 4400MB+-- * Copied: 3MB+-- * Maximum Residency: 0.2MB+-- * Productivity: 99.8%+--+-- These were measured with these constants set:+-- +-- * Participants: 64+-- * Payload Size: 32 words+-- * Total Receives: 3000000+--+-- The memory and productivity numbers do not change much based on+-- the number of capabilities, the nursery size, or the parallel GC+-- settings. In the event-manager-benchmarks repository, I've measured+-- that every call to threadWaitRead causes around 1KB of allocations+-- to happen. This sockets benchmark is certain to call threadWaitRead and+-- unlikely to call threadWaitWrite, so the 3M receives are responsible+-- for 3GB of allocations. It is unclear whether or not improvements+-- to the event manager would result in a tangible gain here.++import Control.Concurrent (forkIO) import Control.Exception (Exception) import Control.Exception (throwIO)-import Control.Monad.ST (runST)-import Data.Primitive (ByteArray)-import Data.Word (Word16,Word8)-import GHC.Exts (RealWorld)-import Test.Tasty-import Test.Tasty.HUnit+import Control.Monad (forever,forM_,when)+import Data.Primitive (PrimArray,MutablePrimArray(..))+import Data.Primitive.MVar (MVar)+import Data.Word (Word16)+import GHC.Clock (getMonotonicTimeNSec)+import GHC.Exts (RealWorld,Int(I#))+import GHC.IO (IO(..))+import Socket.Datagram.IPv4.Undestined (Endpoint(..))+import System.Entropy (getEntropy) -import qualified Socket.Datagram.IPv4.Undestined as DIU-import qualified Socket.Stream.IPv4 as SI-import qualified GHC.Exts as E+import qualified Data.ByteString as B import qualified Data.Primitive as PM import qualified Data.Primitive.MVar as PM+import qualified GHC.Exts as E import qualified Net.IPv4 as IPv4+import qualified Socket.Datagram.IPv4.Undestined as DIU main :: IO () main = do- [duration] <- getArgs- newIORef True- complete <- newPrimArray 1- replicateM_ (take- + done <- PM.newEmptyMVar+ recvCounter <- PM.newPrimArray 1+ PM.writePrimArray recvCounter 0 0+ socketsCounterA <- PM.newPrimArray 1+ PM.writePrimArray socketsCounterA 0 0+ socketsCounterB <- PM.newPrimArray 1+ PM.writePrimArray socketsCounterB 0 0+ socketsA <- PM.newPrimArray participants+ socketsB <- PM.newPrimArray participants+ socketsMVarA <- PM.newEmptyMVar+ socketsMVarB <- PM.newEmptyMVar+ forM_ (enumFromTo 0 (participants - 1)) $ \ix -> do+ forkIO $ worker ix done recvCounter socketsCounterA socketsA socketsMVarA socketsMVarB+ forM_ (enumFromTo 0 (participants - 1)) $ \ix -> do+ forkIO $ worker ix done recvCounter socketsCounterB socketsB socketsMVarB socketsMVarA+ _ <- PM.readMVar socketsMVarA+ _ <- PM.readMVar socketsMVarB+ start <- getMonotonicTimeNSec+ PM.takeMVar done+ end <- getMonotonicTimeNSec+ print (end - start) participants :: Int-participants = 128+participants = 64 +-- This is in units of machine words+payloadSize :: Int+payloadSize = 32+ totalReceives :: Int-totalReceives = 1000000+totalReceives = 3000000 -- The PrimArray must be of length @participants@. worker :: Int -- ^ Worker identifier+ -> MVar RealWorld () -- ^ Used to signal that enough receives have happened+ -> MutablePrimArray RealWorld Int -- ^ Counter of total receives, singleton array -> MutablePrimArray RealWorld Int -- ^ Counter of opened sockets, singleton array- -> MutablePrimArray RealWorld Int -- ^ Counter of total sends, singleton array -> MutablePrimArray RealWorld Word16 -- ^ Ports used by local team- -> MVar (PrimArray Word16) -- ^ MVar for ports used by local team- -> MVar (PrimArray Word16) -- ^ MVar for ports used by remote team+ -> MVar RealWorld (PrimArray Word16) -- ^ MVar for ports used by local team+ -> MVar RealWorld (PrimArray Word16) -- ^ MVar for ports used by remote team -> IO ()-worker !ident !counter !locals !mlocals !mremotes = do- unhandled $ DIU.withSocket (DIU.Endpoint IPv4.loopback 0) $ \sock port -> do- PM.writePrimArray locals ident port- increment counter >>= \case- True -> PM.unsafeFreezePrimArray locals >>= putMVar mlocals+worker !ident !done !recvCounter !counter !locals !mlocals !mremotes = do+ unhandled $ DIU.withSocket (DIU.Endpoint IPv4.loopback 0) $ \sock myPort -> do+ buf@(PM.MutablePrimArray buf#) <- PM.newPrimArray payloadSize+ seedByteString <- getEntropy (payloadSize * PM.sizeOf (undefined :: Int))+ let seedByteArray = E.fromList (B.unpack seedByteString)+ PM.copyByteArray (PM.MutableByteArray buf#) 0 seedByteArray 0 (payloadSize * PM.sizeOf (undefined :: Int))+ PM.writePrimArray locals ident myPort+ incrementWorkerCounter counter >>= \case+ True -> PM.unsafeFreezePrimArray locals >>= PM.putMVar mlocals False -> pure ()- remotes <- readMVar mremotes- act + remotes <- PM.readMVar mremotes+ act sock buf remotes recvCounter done act :: DIU.Socket -- Socket- -> MutableByteArray RealWorld -- Buffer for receives+ -> MutablePrimArray RealWorld Int -- Buffer for receives -> PrimArray Word16 -- Ports used by remote team+ -> MutablePrimArray RealWorld Int -- Receive counter, singleton array+ -> MVar RealWorld () -- Signal that we are finished -> IO ()-act !sock !buf !remotes = case act of- DIU.send sock _ _ _ _- DIU.send sock _ _ _ _+act !sock !buf@(MutablePrimArray buf#) !remotes !counter !done = forever $ do+ n <- scramble buf+ let remote = PM.indexPrimArray remotes (mod n participants)+ unhandled $ DIU.sendMutableByteArray sock (Endpoint {port = remote, address = IPv4.loopback})+ (PM.MutableByteArray buf#) 0 (payloadSize * PM.sizeOf (undefined :: Int))+ recvSz <- unhandled $ DIU.receiveMutableByteArraySlice_ sock (PM.MutableByteArray buf#) 0+ (payloadSize * PM.sizeOf (undefined :: Int))+ when (recvSz /= payloadSize * PM.sizeOf (undefined :: Int)) $ do+ fail "bad receive in act"+ incrementReceiveCounter counter >>= \case+ True -> PM.putMVar done ()+ False -> pure () +scramble :: MutablePrimArray RealWorld Int -> IO Int+scramble arr = go 0 0x36b0b1c47d1ba5e1 0x55109de6a59394b3+ where+ go !ix !acc1 !acc2 = if ix < payloadSize+ then do+ v <- PM.readPrimArray arr ix+ PM.writePrimArray arr ix ((v + acc1) * acc2)+ go (ix + 1) acc2 v+ else pure (acc1 + acc2)+ -- Returns true if the value of the counter reached the total -- number of participants. incrementWorkerCounter :: MutablePrimArray RealWorld Int -> IO Bool-incrementWorkerCounter (MutablePrimArray arr) = IO $ \s0 -> case fetchAddIntArray arr 0# 1# s0 of+incrementWorkerCounter (MutablePrimArray arr) = IO $ \s0 -> case E.fetchAddIntArray# arr 0# 1# s0 of (# s1, i #) -> (# s1, I# i == participants - 1 #) incrementReceiveCounter :: MutablePrimArray RealWorld Int -> IO Bool-incrementReceiveCounter (MutablePrimArray arr) = IO $ \s0 -> case fetchAddIntArray arr 0# 1# s0 of+incrementReceiveCounter (MutablePrimArray arr) = IO $ \s0 -> case E.fetchAddIntArray# arr 0# 1# s0 of (# s1, i #) -> (# s1, I# i == totalReceives - 1 #) -tests :: TestTree-tests = testGroup "socket"- [ testGroup "datagram"- [ testGroup "ipv4"- [ testGroup "undestined"- [ testCase "A" testDatagramUndestinedA- ]- ]- ]- , testGroup "stream"- [ testGroup "ipv4"- [ testCase "A" testStreamA- ]- ]- ]- unhandled :: Exception e => IO (Either e a) -> IO a unhandled action = action >>= either throwIO pure--testDatagramUndestinedA :: Assertion-testDatagramUndestinedA = do- (m :: PM.MVar RealWorld Word16) <- PM.newEmptyMVar- (port,received) <- concurrently (sender m) (receiver m)- received @=? (DIU.Endpoint IPv4.loopback port, message)- where- message = E.fromList [0,1,2,3] :: ByteArray- sz = PM.sizeofByteArray message- sender :: PM.MVar RealWorld Word16 -> IO Word16- sender m = unhandled $ DIU.withSocket (DIU.Endpoint IPv4.loopback 0) $ \sock srcPort -> do- dstPort <- PM.takeMVar m- unhandled $ DIU.send sock (DIU.Endpoint IPv4.loopback dstPort) message 0 sz- pure srcPort- receiver :: PM.MVar RealWorld Word16 -> IO (DIU.Endpoint,ByteArray)- receiver m = unhandled $ DIU.withSocket (DIU.Endpoint IPv4.loopback 0) $ \sock port -> do- PM.putMVar m port- unhandled $ DIU.receive sock sz---- This test involves a made up protocol that goes like this:--- The sender always starts by sending the length of the rest--- of the payload as a native-endian encoded machine-sized int.--- (This could only ever work for a machine that is communicating--- with itself). Then, it sends a bytearray of that specified--- length. Then, both ends are expected to shutdown their sides--- of the connection.-testStreamA :: Assertion-testStreamA = do- (m :: PM.MVar RealWorld Word16) <- PM.newEmptyMVar- ((),received) <- concurrently (sender m) (receiver m)- received @=? message- where- message = E.fromList (enumFromTo 0 (100 :: Word8)) :: ByteArray- sz = PM.sizeofByteArray message- szb = runST $ do- marr <- PM.newByteArray (PM.sizeOf (undefined :: Int))- PM.writeByteArray marr 0 sz- PM.unsafeFreezeByteArray marr- sender :: PM.MVar RealWorld Word16 -> IO ()- sender m = do- dstPort <- PM.takeMVar m- unhandled $ SI.withConnection (DIU.Endpoint IPv4.loopback dstPort) $ \conn -> do- unhandled $ SI.sendByteArray conn szb- unhandled $ SI.sendByteArray conn message- receiver :: PM.MVar RealWorld Word16 -> IO ByteArray- receiver m = unhandled $ SI.withListener (SI.Endpoint IPv4.loopback 0) $ \listener port -> do- PM.putMVar m port- unhandled $ SI.withAccepted listener $ \conn _ -> do- serializedSize <- unhandled $ SI.receiveByteArray conn (PM.sizeOf (undefined :: Int))- let theSize = PM.indexByteArray serializedSize 0 :: Int- result <- unhandled $ SI.receiveByteArray conn theSize- pure result--
example/Main.hs view
@@ -53,7 +53,7 @@ unhandled $ DIU.withSocket (DIU.Endpoint IPv4.loopback 0) $ \sock port -> do BC.putStrLn ("Receiving datagrams on 127.0.0.1:" <> BC.pack (show port)) replicateM_ 10 $ do- (remote,ByteArray payload) <- unhandled (DIU.receive sock 1024)+ DIU.Message remote (ByteArray payload) <- unhandled (DIU.receive sock 1024) BC.putStrLn ("Datagram from " <> BC.pack (show remote)) BC.putStr (SB.fromShort (SB.SBS payload))
sockets.cabal view
@@ -1,8 +1,17 @@ cabal-version: 2.2 name: sockets-version: 0.1.0.0+version: 0.3.0.0 synopsis: High-level network sockets-description: High-level abstraction for network sockets+description:+ This library provides a high-level abstraction for network sockets. It uses+ Haskell2010 (along with GADTs) without typeclasses to ensure that+ consumers of the API can only call appropriate functions on a socket.+ Exceptions are tracked in the types of functions and returned to the caller+ with `Either`. The caller is free to handle these gracefully or to throw+ them. This library only throws exceptions when it detects that it has misused+ the operating system's sockets API (open an issue for this) or when the+ caller asks for a negatively-sized slice of a buffer (such exceptions are+ unrecoverable and indicate a mistake in the code consuming this API). homepage: https://github.com/andrewthad/sockets bug-reports: https://github.com/andrewthad/sockets/issues license: BSD-3-Clause@@ -13,6 +22,11 @@ category: Network extra-source-files: CHANGELOG.md +flag mmsg+ manual: True+ description: Use sendmmsg and recvmmsg + default: False+ flag debug manual: True description: Print debug output @@ -26,21 +40,32 @@ library exposed-modules: Socket.Datagram.IPv4.Undestined+ Socket.Datagram.IPv4.Spoof Socket.Stream.IPv4 other-modules:+ Socket.Stream+ Socket.Datagram+ Socket.Datagram.IPv4.Undestined.Multiple+ Socket.Datagram.IPv4.Undestined.Internal Socket.Debug Socket.IPv4 Socket build-depends: , base >= 4.11.1.0 && < 5- , posix-api >= 0.2- , primitive >= 0.6.4 , ip >= 1.4.1+ , posix-api >= 0.2.1+ , primitive >= 0.6.4+ , stm >= 2.4+ , text >= 1.2 hs-source-dirs: src if flag(debug) hs-source-dirs: src-debug else hs-source-dirs: src-production+ if flag(mmsg)+ hs-source-dirs: src-mmsg+ else+ hs-source-dirs: src-no-mmsg default-language: Haskell2010 ghc-options: -O2 -Wall @@ -66,7 +91,9 @@ , sockets , ip >= 1.4.1 , primitive >= 0.6.4- ghc-options: -Wall -O2 -threaded+ , bytestring >= 0.10.8.2+ , entropy >= 0.4.1.4+ ghc-options: -Wall -O2 -threaded -rtsopts default-language: Haskell2010 hs-source-dirs: bench main-is: Macro.hs
src-debug/Socket/Debug.hs view
@@ -1,6 +1,10 @@ module Socket.Debug ( debug+ , whenDebugging ) where debug :: String -> IO () debug = putStrLn++whenDebugging :: IO () -> IO ()+whenDebugging = id
+ src-mmsg/Socket/Datagram/IPv4/Undestined/Multiple.hs view
@@ -0,0 +1,101 @@+{-# language BangPatterns #-}+{-# language DuplicateRecordFields #-}+{-# language LambdaCase #-}+{-# language MagicHash #-}+{-# language NamedFieldPuns #-}+{-# language UnboxedTuples #-}+module Socket.Datagram.IPv4.Undestined.Multiple+ ( receiveMany+ ) where++import Control.Concurrent (threadWaitWrite,threadWaitRead)+import Control.Exception (mask,onException)+import Data.Primitive (ByteArray,MutableByteArray(..),Array)+import Data.Word (Word16)+import Foreign.C.Error (Errno(..),eWOULDBLOCK,eAGAIN)+import Foreign.C.Types (CInt,CSize,CUInt)+import GHC.Exts (Int(I#),RealWorld,shrinkMutableByteArray#,ByteArray#,touch#)+import GHC.IO (IO(..))+import Net.Types (IPv4(..))+import Socket (SocketException(..))+import Socket.Datagram.IPv4.Undestined.Internal (Message(..),Socket(..))+import Socket.Debug (debug)+import Socket.IPv4 (Endpoint(..))+import System.Posix.Types (Fd)++import qualified Control.Monad.Primitive as PM+import qualified Data.Primitive as PM+import qualified Linux.Socket as L+import qualified Posix.Socket as S++-- | Receive up to the specified number of datagrams into freshly allocated+-- byte arrays. When there are many datagrams present on the receive+-- buffer, this is more efficient than calling 'receive' repeatedly. The+-- array is guaranteed to have at least one message.+--+-- The byte arrays in the resulting messages are always pinned.+receiveMany ::+ Socket -- ^ Socket+ -> Int -- ^ Maximum number of datagrams to receive+ -> Int -- ^ Maximum size of each datagram to receive+ -> IO (Either SocketException (Array Message))+receiveMany = receiveManyNative++receiveManyNative :: Socket -> Int -> Int -> IO (Either SocketException (Array Message))+receiveManyNative (Socket !fd) !maxDatagrams !maxSz = do+ threadWaitRead fd+ L.uninterruptibleReceiveMultipleMessageB fd S.sizeofSocketAddressInternet (intToCSize maxSz) (intToCUInt maxDatagrams) L.truncate >>= \case+ Left err -> pure (Left (errorCode err))+ Right (saneSockAddrs,sockAddrs,greatestMsgSz,msgs) -> if saneSockAddrs == 0+ then if cuintToInt greatestMsgSz > maxSz+ then pure (Left (ReceivedMessageTruncated (cuintToInt greatestMsgSz)))+ else do+ let len = PM.sizeofUnliftedArray msgs+ let sockaddrBase = PM.byteArrayContents sockAddrs+ finalMsgs <- PM.newArray len errorThunk+ let go !ix = if ix >= 0+ then S.indexSocketAddressInternet sockaddrBase ix >>= \case+ Left fam -> do+ touchByteArray sockAddrs+ pure (Left (SocketAddressFamily fam))+ Right sockAddrInet -> do+ touchByteArray sockAddrs+ let !msg = Message+ (socketAddressInternetToEndpoint sockAddrInet)+ (PM.indexUnliftedArray msgs ix)+ PM.writeArray finalMsgs ix msg+ go (ix - 1)+ else do+ touchByteArray sockAddrs+ fmap Right (PM.unsafeFreezeArray finalMsgs)+ go (len - 1)+ else pure (Left SocketAddressSize)++-- Used internally in arrays+errorThunk :: a+errorThunk = error "Socket.Datagram.IPv4.Undestined: uninitialized element"++touchByteArray :: ByteArray -> IO ()+touchByteArray (PM.ByteArray x) = touchByteArray# x++touchByteArray# :: ByteArray# -> IO ()+touchByteArray# x = IO $ \s -> case touch# x s of s' -> (# s', () #)++socketAddressInternetToEndpoint :: S.SocketAddressInternet -> Endpoint+socketAddressInternetToEndpoint (S.SocketAddressInternet {address,port}) = Endpoint+ { address = IPv4 (S.networkToHostLong address)+ , port = S.networkToHostShort port+ }++cuintToInt :: CUInt -> Int+cuintToInt = fromIntegral++errorCode :: Errno -> SocketException+errorCode (Errno x) = ErrorCode x++intToCUInt :: Int -> CUInt+intToCUInt = fromIntegral++intToCSize :: Int -> CSize+intToCSize = fromIntegral+
+ src-no-mmsg/Socket/Datagram/IPv4/Undestined/Multiple.hs view
@@ -0,0 +1,140 @@+{-# language BangPatterns #-}+{-# language DuplicateRecordFields #-}+{-# language LambdaCase #-}+{-# language MagicHash #-}+{-# language NamedFieldPuns #-}+{-# language UnboxedTuples #-}+module Socket.Datagram.IPv4.Undestined.Multiple+ ( receiveMany+ , receiveManyUnless+ ) where++import Control.Applicative ((<|>))+import Control.Monad.STM (STM,atomically)+import Control.Concurrent (threadWaitWrite,threadWaitRead,threadWaitReadSTM)+import Control.Exception (mask,onException)+import Data.Functor (($>))+import Data.Primitive (ByteArray,MutableByteArray(..),Array)+import Data.Word (Word16)+import Foreign.C.Error (Errno(..),eWOULDBLOCK,eAGAIN)+import Foreign.C.Types (CInt,CSize,CUInt)+import GHC.Exts (Int(I#),RealWorld,shrinkMutableByteArray#,ByteArray#,touch#)+import GHC.IO (IO(..))+import Net.Types (IPv4(..))+import Socket (SocketException(..))+import Socket.Datagram.IPv4.Undestined.Internal (Message(..),Socket(..))+import Socket.Debug (debug)+import Socket.IPv4 (Endpoint(..))+import System.Posix.Types (Fd)++import qualified Control.Monad.Primitive as PM+import qualified Data.Primitive as PM+import qualified Linux.Socket as L+import qualified Posix.Socket as S++-- | Receive up to the specified number of datagrams into freshly allocated+-- byte arrays. When there are many datagrams present on the receive+-- buffer, this is more efficient than calling 'receive' repeatedly. The+-- array is guaranteed to have at least one message.+--+-- The byte arrays in the resulting messages are always pinned.+receiveMany ::+ Socket -- ^ Socket+ -> Int -- ^ Maximum number of datagrams to receive+ -> Int -- ^ Maximum size of each datagram to receive+ -> IO (Either SocketException (Array Message))+receiveMany (Socket !fd) !maxDatagrams !maxSz = do+ debug "receiveMany: about to wait"+ threadWaitRead fd+ receiveManyShim fd maxDatagrams maxSz++-- | This has the same behavior as 'receiveMany'. However, it also takes an+-- 'STM' action that it attempts to run while the event manager is waiting+-- for the socket to be ready for a reads. If the supplied action finishes+-- first, this abandons the attempt to receive datagrams and returns+-- @'Left' 'ReceptionAbandoned'@.+receiveManyUnless :: + STM () -- ^ If this completes, give up on receiving+ -> Socket -- ^ Socket+ -> Int -- ^ Maximum number of datagrams to receive+ -> Int -- ^ Maximum size of each datagram to receive+ -> IO (Either SocketException (Array Message))+receiveManyUnless abandon (Socket !fd) !maxDatagrams !maxSz = do+ debug "receiveMany: about to wait"+ (isReady,deregister) <- threadWaitReadSTM fd+ shouldReceive <- atomically ((abandon $> False) <|> (isReady $> True))+ deregister+ if shouldReceive+ then receiveManyShim fd maxDatagrams maxSz+ else pure (Left ReceptionAbandoned)++-- Although this is a shim for recvmmsg, it is still better than calling+-- receive repeatedly since it avoids unneeded calls to the event+-- manager. This is guaranteed to return at least one message.+--+-- This function is currently unused. It is being left here so that,+-- when cross-platform compatibility is someday handled, this will+-- be available for platforms that do not provide recvmmsg.+receiveManyShim :: Fd -> Int -> Int -> IO (Either SocketException (Array Message))+receiveManyShim !fd !maxDatagrams !maxSz = do+ debug "receiveMany: socket is now readable"+ msgs <- PM.newArray maxDatagrams errorThunk+ -- We use MSG_TRUNC so that we are able to figure out whether+ -- or not bytes were discarded. If bytes were discarded+ -- (meaning that the buffer was too small), we return an+ -- exception.+ let go !ix = if ix < maxDatagrams+ then do+ -- This does not need to allocate pinned memory for+ -- the call to recvfrom to work correctly. It allocates+ -- pinned memory so that its behavior is consistent with+ -- that of receiveManyNative.+ marr <- PM.newPinnedByteArray maxSz+ e <- S.uninterruptibleReceiveFromMutableByteArray fd marr 0+ (intToCSize maxSz) (L.truncate) S.sizeofSocketAddressInternet+ case e of+ Left err -> if err == eWOULDBLOCK || err == eAGAIN+ then do+ r <- PM.freezeArray msgs 0 ix+ pure (Right r)+ else pure (Left (errorCode err))+ Right (sockAddrRequiredSz,sockAddr,recvSz) -> if csizeToInt recvSz <= maxSz+ then if sockAddrRequiredSz == S.sizeofSocketAddressInternet+ then case S.decodeSocketAddressInternet sockAddr of+ Just sockAddrInet -> do+ shrinkMutableByteArray marr (csizeToInt recvSz)+ arr <- PM.unsafeFreezeByteArray marr+ let !msg = Message (socketAddressInternetToEndpoint sockAddrInet) arr+ PM.writeArray msgs ix msg+ go (ix + 1)+ Nothing -> pure (Left (SocketAddressFamily (-1)))+ else pure (Left SocketAddressSize)+ else pure (Left (ReceivedMessageTruncated (csizeToInt recvSz)))+ else do+ r <- PM.unsafeFreezeArray msgs+ pure (Right r)+ go 0++-- Used internally in arrays+errorThunk :: a+errorThunk = error "Socket.Datagram.IPv4.Undestined: uninitialized element"++csizeToInt :: CSize -> Int+csizeToInt = fromIntegral++socketAddressInternetToEndpoint :: S.SocketAddressInternet -> Endpoint+socketAddressInternetToEndpoint (S.SocketAddressInternet {address,port}) = Endpoint+ { address = IPv4 (S.networkToHostLong address)+ , port = S.networkToHostShort port+ }++shrinkMutableByteArray :: MutableByteArray RealWorld -> Int -> IO ()+shrinkMutableByteArray (MutableByteArray arr) (I# sz) =+ PM.primitive_ (shrinkMutableByteArray# arr sz)++intToCSize :: Int -> CSize+intToCSize = fromIntegral++errorCode :: Errno -> SocketException+errorCode (Errno x) = ErrorCode x+
src-production/Socket/Debug.hs view
@@ -1,6 +1,10 @@ module Socket.Debug ( debug+ , whenDebugging ) where debug :: String -> IO () debug _ = pure ()++whenDebugging :: IO () -> IO ()+whenDebugging _ = pure ()
src/Socket.hs view
@@ -1,54 +1,65 @@ {-# language BangPatterns #-}+{-# language DataKinds #-} {-# language DeriveAnyClass #-} {-# language DerivingStrategies #-} {-# language DuplicateRecordFields #-}+{-# language GADTs #-}+{-# language KindSignatures #-} module Socket ( SocketException(..)- , Context(..)- , Reason(..)+ , SocketUnrecoverableException(..)+ , Direction(..)+ , Interruptibility(..)+ , Forkedness(..)+ , cgetsockname + , cgetsockopt+ , cclose + , crecv + , crecvfrom+ , cshutdown+ , negativeSliceLength+ , nonInternetSocketFamily+ , functionWithAccepted+ , functionWithConnection+ , functionWithListener+ , functionWithSocket+ , functionGracefulClose+ , socketAddressSize ) where -import Control.Exception (Exception)+import Control.Exception (Exception(..))+import Data.Kind (Type) import Foreign.C.Types (CInt) --- | Represents any unexpected behaviors that a function working on a--- socket, connection, or listener can exhibit.-data SocketException = SocketException- { context :: Context- , reason :: Reason- }- deriving stock (Eq,Show)- deriving anyclass (Exception)+import qualified Data.List as L --- | The function that behaved unexpectedly.-data Context- = Accept- | Bind- | Close- | Connect- | GetName- | Listen- | Open- | Option- | Receive- | Send- | Shutdown- deriving stock (Eq,Show)+data Direction = Send | Receive --- | A description of the unexpected behavior.-data Reason- = MessageTruncated !Int !Int+data Interruptibility = Interruptible | Uninterruptible++data Forkedness = Forked | Unforked++-- | Represents any unexpected behaviors that a function working on a+-- socket, connection, or listener can exhibit.+data SocketException+ = SentMessageTruncated !Int -- ^ The datagram did not fit in the buffer. This can happen while- -- sending or receiving. Fields: buffer size, datagram size.+ -- sending. The field is the size of the number of bytes in the+ -- datagram that were successfully copied into the send buffer.+ | ReceivedMessageTruncated !Int+ -- ^ The datagram did not fit in the buffer. This can happen while+ -- receiving. The field is the original size of the datagram that+ -- was was truncated while copying it into the buffer. | SocketAddressSize -- ^ The socket address was not the expected size. This exception -- indicates a bug in this library or (less likely) in the -- operating system.- | SocketAddressFamily+ | SocketAddressFamily !CInt -- ^ The socket address had an unexpected family. This exception -- indicates a bug in this library or (less likely) in the- -- operating system.+ -- operating system. The int argument is the actual family+ -- found in the socket address. | OptionValueSize -- ^ The option value was not the expected size. This exception -- indicates a bug in this library or (less likely) in the@@ -56,6 +67,10 @@ | NegativeBytesRequested -- ^ The user requested a negative number of bytes in a call -- to a receive function.+ | ReceptionAbandoned+ -- ^ This happens when the @Unless@ variant of a function is+ -- used and the @STM@ action completes before the socket is+ -- ready for a read. | RemoteNotShutdown -- ^ The remote end sent more data when it was expected to send -- a shutdown.@@ -72,3 +87,57 @@ -- not expect or recognize. Consult your operating system manual -- for details about the error code. deriving stock (Eq,Show)+ deriving anyclass (Exception)++data SocketUnrecoverableException = SocketUnrecoverableException+ { modules :: String+ , function :: String+ , description :: [String]+ }+ deriving stock (Show,Eq)++instance Exception SocketUnrecoverableException where+ displayException (SocketUnrecoverableException m f d) =+ m ++ "." ++ f ++ ": [" ++ L.intercalate "," d ++ "]"++cgetsockname :: String+cgetsockname = "getsockname"++cgetsockopt :: String+cgetsockopt = "getsockopt"++cclose :: String+cclose = "getsockname"++crecv :: String+crecv = "recv"++crecvfrom :: String+crecvfrom = "recvfrom"++cshutdown :: String+cshutdown = "shutdown"++functionGracefulClose :: String+functionGracefulClose = "gracefulClose"++nonInternetSocketFamily :: String+nonInternetSocketFamily = "non-internet socket family"++negativeSliceLength :: String+negativeSliceLength = "negative slice length"++functionWithAccepted :: String+functionWithAccepted = "withAccepted"++functionWithConnection :: String+functionWithConnection = "withConnection"++functionWithListener :: String+functionWithListener = "withListener"++functionWithSocket :: String+functionWithSocket = "withSocket"++socketAddressSize :: String+socketAddressSize = "socket address size"
+ src/Socket/Datagram.hs view
@@ -0,0 +1,44 @@+{-# language DataKinds #-}+{-# language DeriveAnyClass #-}+{-# language DerivingStrategies #-}+{-# language GADTs #-}+{-# language KindSignatures #-}+{-# language StandaloneDeriving #-}+module Socket.Datagram+ ( SendException(..)+ , ReceiveException(..)+ , SocketException(..)+ ) where++import Socket (Interruptibility(..))+import Socket.IPv4 (SocketException(..))++import Data.Kind (Type)+import Data.Typeable (Typeable)+import Control.Exception (Exception)++data SendException :: Interruptibility -> Type where+ -- | The datagram did not fit in the buffer. The field is the+ -- number of bytes that were successfully copied into the+ -- send buffer. The datagram does still get sent when this+ -- happens.+ SendTruncated :: !Int -> SendException i+ -- | Attempted to send to a broadcast address.+ SendBroadcasted :: SendException i+ -- | STM-style interrupt (much safer than C-style interrupt)+ SendInterrupted :: SendException 'Interruptible++deriving stock instance Show (SendException i)+deriving anyclass instance (Typeable i) => Exception (SendException i)++data ReceiveException :: Interruptibility -> Type where+ -- | The datagram did not fit in the buffer. The field is the+ -- original size of the datagram that was truncated. If+ -- this happens, the process probably needs to start using+ -- a larger receive buffer.+ ReceiveTruncated :: !Int -> ReceiveException i+ -- | STM-style interrupt (much safer than C-style interrupt)+ ReceiveInterrupted :: ReceiveException 'Interruptible++deriving stock instance Show (ReceiveException i)+deriving anyclass instance (Typeable i) => Exception (ReceiveException i)
+ src/Socket/Datagram/IPv4/Spoof.hs view
@@ -0,0 +1,290 @@+{-# language BangPatterns #-}+{-# language DataKinds #-}+{-# language DeriveAnyClass #-}+{-# language DerivingStrategies #-}+{-# language DuplicateRecordFields #-}+{-# language GADTs #-}+{-# language KindSignatures #-}+{-# language LambdaCase #-}+{-# language MagicHash #-}+{-# language NamedFieldPuns #-}+{-# language StandaloneDeriving #-}+{-# language UnboxedTuples #-}++-- | Internet datagram sockets without a fixed destination.+-- The user may spoof the source address and may specify the+-- packet ID. An application must have @CAP_NET_RAW@ or be+-- running as root to use the functions in this module.+module Socket.Datagram.IPv4.Spoof+ ( -- * Types+ Socket(..)+ , Endpoint(..)+ , Message(..)+ -- * Establish+ , withSocket+ -- * Communicate+ , sendMutableByteArray+ -- * Exceptions+ , SocketException(..)+ , SendException(..)+ ) where++import Control.Concurrent (threadWaitWrite)+import Control.Exception (Exception,throwIO,mask,onException)+import Data.Bits (unsafeShiftR,complement,(.&.))+import Data.Kind (Type)+import Data.Primitive (MutableByteArray(..))+import Data.Word (Word16,Word8,Word64,Word32)+import Foreign.C.Error (Errno(..),eWOULDBLOCK,eAGAIN,eMFILE,eNFILE,eACCES,ePERM)+import Foreign.C.Types (CInt,CSize)+import GHC.Exts (RealWorld,touch#)+import GHC.IO (IO(..))+import Net.Types (IPv4(..))+import Socket (SocketUnrecoverableException(..),Interruptibility(..))+import Socket.Datagram (SendException(..))+import Socket.Datagram.IPv4.Undestined.Internal (Message(..))+import Socket.Debug (debug,whenDebugging)+import Socket.IPv4 (Endpoint(..))+import System.Posix.Types (Fd)+import Text.Printf (printf)++import qualified Data.Primitive as PM+import qualified Linux.Socket as L+import qualified Posix.Socket as S+import qualified GHC.Exts as E+import qualified Socket as SCK++-- TODO: Something I am not sure about is whether or not it is necessary+-- to bind to a port right after creating the socket. If we defering+-- binding until the call to the time of sending, what port do we bind+-- to? Does the kernel use the one in the source port or does it choose+-- an ephemeral port? We need it to choose an ephemeral port. Otherwise,+-- we can get spurious failures.++-- | A socket that send datagrams with spoofed source IP addresses.+-- It cannot receive datagrams.+newtype Socket = Socket Fd+ deriving stock (Eq,Ord,Show)++data SocketException :: Type where+ -- | Permission to create a raw socket was denied. The process needs+ -- the capability @CAP_NET_RAW@, or it must be run as root.+ SocketPermissionDenied :: SocketException+ -- | A limit on the number of open file descriptors has been reached.+ -- This could be the per-process limit or the system limit.+ -- (@EMFILE@ and @ENFILE@)+ SocketFileDescriptorLimit :: SocketException++deriving stock instance Show SocketException+deriving anyclass instance Exception SocketException++-- | Open a socket and run the supplied callback on it. This closes the socket+-- when the callback finishes or when an exception is thrown. Do not return +-- the socket from the callback. This leads to undefined behavior. The user+-- cannot specify an endpoint since the socket cannot receive traffic.+withSocket ::+ (Socket -> IO a) -- ^ Callback providing the socket+ -> IO (Either SocketException a)+withSocket f = mask $ \restore -> do+ debug "withSocket: opening raw socket"+ e1 <- S.uninterruptibleSocket S.internet+ (L.applySocketFlags (L.closeOnExec <> L.nonblocking) S.raw)+ S.rawProtocol+ debug "withSocket: opened raw socket"+ case e1 of+ Left err -> handleSocketException SCK.functionWithSocket err+ Right fd -> do+ a <- onException (restore (f (Socket fd))) (S.uninterruptibleErrorlessClose fd)+ S.uninterruptibleClose fd >>= \case+ Left err -> throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Spoof+ SCK.functionWithSocket+ ["close",describeErrorCode err]+ Right _ -> pure (Right a)++-- | Send a slice of a bytearray to the specified endpoint.+sendMutableByteArray ::+ Socket -- ^ Socket+ -> Endpoint -- ^ Spoofed source address and port+ -> Endpoint -- ^ Remote IPv4 address and port+ -> MutableByteArray RealWorld -- ^ Buffer (will be sliced)+ -> Int -- ^ Offset into payload+ -> Int -- ^ Lenth of slice into buffer+ -> IO (Either (SendException 'Uninterruptible) ())+sendMutableByteArray (Socket !s) !theSource !theRemote !thePayload !off !len = do+ let ipHeaderSz = cintToInt L.sizeofIpHeader+ let totalHeaderSz = cintToInt (L.sizeofIpHeader + L.sizeofUdpHeader)+ let totalPacketSz = len + totalHeaderSz+ -- Why do we add one to the size? This extra byte at the end will always+ -- be zeroed out. It makes UDP checksum calculation a little easier,+ -- since we can now pull out Word16s until we reach the end. If the+ -- original length was even, this extra byte ends up unused by the+ -- checksum. But, if the original length was odd, it does get used.+ buf <- PM.newPinnedByteArray (totalPacketSz + 1)+ PM.setByteArray buf 0 (totalPacketSz + 1) (0 :: Word8)+ let addr = PM.mutableByteArrayContents buf+ L.pokeIpHeaderVersionIhl addr (4 * 16 + 5)+ L.pokeIpHeaderTypeOfService addr 0+ -- TODO: Actually check the length.+ -- NB: The packet length must be in network byte order.+ -- Expermentally, it seems that it does not.+ L.pokeIpHeaderTotalLength addr (S.hostToNetworkShort (intToWord16 totalPacketSz))+ L.pokeIpHeaderIdentifier addr 0+ L.pokeIpHeaderFragmentOffset addr 0+ L.pokeIpHeaderTimeToLive addr 64+ L.pokeIpHeaderProtocol addr (cintToWord8 (S.getProtocol S.udp))+ -- The linux kernel fills in the ip header checksum for us.+ L.pokeIpHeaderChecksum addr 0+ let src = S.hostToNetworkLong (getIPv4 (address theSource))+ L.pokeIpHeaderSourceAddress addr src+ let dst = S.hostToNetworkLong (getIPv4 (address theRemote))+ L.pokeIpHeaderDestinationAddress addr dst+ let udpAddr = PM.plusAddr addr ipHeaderSz+ L.pokeUdpHeaderSourcePort udpAddr (S.hostToNetworkShort (port theSource))+ L.pokeUdpHeaderDestinationPort udpAddr (S.hostToNetworkShort (port theRemote))+ let udpLen = cintToInt L.sizeofUdpHeader + len+ L.pokeUdpHeaderLength udpAddr (S.hostToNetworkShort (intToWord16 udpLen))+ PM.copyMutableByteArray buf totalHeaderSz thePayload off len+ L.pokeUdpHeaderChecksum udpAddr . S.hostToNetworkShort =<< udpChecksum src dst buf ipHeaderSz udpLen+ touchMutableByteArray buf+ debug ("spoof send mutable: about to send to " ++ show theRemote)+ whenDebugging $ do+ d <- PM.newByteArray totalPacketSz+ PM.copyMutableByteArray d 0 buf 0 totalPacketSz+ x <- PM.unsafeFreezeByteArray d+ debug ("raw packet: " ++ (foldMap (printf "%.2x ") (E.toList x)))+ e1 <- S.uninterruptibleSendToMutableByteArray s buf 0 (intToCSize totalPacketSz)+ mempty+ (S.encodeSocketAddressInternet (endpointToSocketAddressInternet theRemote))+ debug ("spoof send mutable: just sent to " ++ show theRemote)+ case e1 of+ Left err1 -> if err1 == eWOULDBLOCK || err1 == eAGAIN+ then do+ debug ("send mutable: waiting to for write ready to send to " ++ show theRemote)+ threadWaitWrite s+ e2 <- S.uninterruptibleSendToMutableByteArray s buf+ (intToCInt off)+ (intToCSize len)+ mempty+ (S.encodeSocketAddressInternet (endpointToSocketAddressInternet theRemote))+ case e2 of+ Left err2 -> do+ debug ("send mutable: encountered error after sending")+ handleSendException "sendMutableByteArray" err2+ Right sz -> if csizeToInt sz == totalPacketSz+ then pure (Right ())+ else pure (Left (SendTruncated (csizeToInt sz)))+ else do+ debug "spoof send mutable: sent on first try but got error code" + handleSendException "sendMutableByteArray" err1+ Right sz -> if csizeToInt sz == totalPacketSz+ then do+ debug ("send mutable: success")+ pure (Right ())+ else pure (Left (SendTruncated (csizeToInt sz)))++-- Precondition: the mutable byte array must have an extra zeroed out+-- byte at the end. That is, at arr[offset+length], there exists a+-- zero byte. The offset must divide two evenly.+udpChecksum ::+ Word32 -- source (network byte order)+ -> Word32 -- dest (network byte order)+ -> MutableByteArray RealWorld -- payload+ -> Int -- offset (start of the udp header)+ -> Int -- length (udp header size plus payload size) + -> IO Word16+udpChecksum src dst payload off len = do+ let sum0 = word16ToWord64 (S.hostToNetworkShort (word32ToWord16 src))+ debug ("udp checksum source lower: " ++ printf "%.8X" sum0)+ let sum1 = sum0 + word16ToWord64 (S.hostToNetworkShort (word32ToWord16 (unsafeShiftR src 16)))+ debug ("udp checksum source lower+upper: " ++ printf "%.8X" sum1)+ let sum2 = sum1 + word16ToWord64 (S.hostToNetworkShort (word32ToWord16 dst))+ sum3 = sum2 + word16ToWord64 (S.hostToNetworkShort (word32ToWord16 (unsafeShiftR dst 16)))+ debug ("udp checksum source+dest lower+upper: " ++ printf "%.8X" sum3)+ let sum4 = sum3 + word16ToWord64 (cintToWord16 (S.getProtocol S.udp))+ sum5 = sum4 + word16ToWord64 (intToWord16 len)+ debug ("udp checksum pseudoheader without carries: " ++ printf "%.8X" sum5)+ let halfLen = unsafeShiftR (len + off) 1+ debug ("udp checksum start offset: " ++ show (unsafeShiftR off 1))+ debug ("udp checksum last offset: " ++ show halfLen)+ let go :: Int -> Word64 -> IO Word64+ go !ix !acc = if ix < halfLen+ then do+ w16 <- PM.readByteArray payload ix :: IO Word16+ debug ("udp checksum payload iteration " ++ show ix ++ ": " ++ printf "%.8X" acc)+ go (ix + 1) (word16ToWord64 (S.hostToNetworkShort w16) + acc)+ else pure acc+ r <- go (unsafeShiftR off 1) sum5+ -- We assume that the upper 16 bits in this 64-bit word are zeroes.+ -- There is no way for a datagram to be long enough to start to+ -- fill the bits beyond 48.+ pure (word64ToWord16 (complement ((r .&. 0xFFFF) + (unsafeShiftR r 16 .&. 0xFFFF) + (unsafeShiftR r 32))))++endpointToSocketAddressInternet :: Endpoint -> S.SocketAddressInternet+endpointToSocketAddressInternet (Endpoint {address, port}) = S.SocketAddressInternet+ { port = S.hostToNetworkShort port+ , address = S.hostToNetworkLong (getIPv4 address)+ }++intToCInt :: Int -> CInt+intToCInt = fromIntegral++intToCSize :: Int -> CSize+intToCSize = fromIntegral++csizeToInt :: CSize -> Int+csizeToInt = fromIntegral++cintToInt :: CInt -> Int+cintToInt = fromIntegral++cintToWord8 :: CInt -> Word8+cintToWord8 = fromIntegral++intToWord16 :: Int -> Word16+intToWord16 = fromIntegral++cintToWord16 :: CInt -> Word16+cintToWord16 = fromIntegral++word16ToWord64 :: Word16 -> Word64+word16ToWord64 = fromIntegral++word64ToWord16 :: Word64 -> Word16+word64ToWord16 = fromIntegral++word32ToWord16 :: Word32 -> Word16+word32ToWord16 = fromIntegral++touchMutableByteArray :: MutableByteArray RealWorld -> IO ()+touchMutableByteArray (MutableByteArray x) = touchMutableByteArray# x++touchMutableByteArray# :: E.MutableByteArray# RealWorld -> IO ()+touchMutableByteArray# x = IO $ \s -> case touch# x s of s' -> (# s', () #)++moduleSocketDatagramIPv4Spoof :: String+moduleSocketDatagramIPv4Spoof = "Socket.Datagram.IPv4.Spoof"++handleSocketException :: String -> Errno -> IO (Either SocketException a)+{-# INLINE handleSocketException #-}+handleSocketException func e+ | e == ePERM = pure (Left SocketPermissionDenied)+ | e == eMFILE = pure (Left SocketFileDescriptorLimit)+ | e == eNFILE = pure (Left SocketFileDescriptorLimit)+ | otherwise = throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Spoof+ func+ [describeErrorCode e]++describeErrorCode :: Errno -> String+describeErrorCode (Errno e) = "error code " ++ show e++handleSendException :: String -> Errno -> IO (Either (SendException i) a)+{-# INLINE handleSendException #-}+handleSendException func e+ | e == eACCES = pure (Left SendBroadcasted)+ | otherwise = throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Spoof+ func+ [describeErrorCode e]+
src/Socket/Datagram/IPv4/Undestined.hs view
@@ -1,52 +1,57 @@ {-# language BangPatterns #-}+{-# language DataKinds #-} {-# language DeriveAnyClass #-} {-# language DerivingStrategies #-} {-# language DuplicateRecordFields #-} {-# language LambdaCase #-} {-# language MagicHash #-} {-# language NamedFieldPuns #-}+{-# language UnboxedTuples #-} +-- | Internet datagram sockets without a fixed destination. module Socket.Datagram.IPv4.Undestined ( -- * Types Socket(..) , Endpoint(..)+ , Message(..) -- * Establish , withSocket -- * Communicate , send- , receive+ , sendMutableByteArraySlice+ , receiveByteArray , receiveMutableByteArraySlice_+ , receiveMany+ , receiveManyUnless -- * Exceptions , SocketException(..)- , Context(..)- , Reason(..) -- * Examples -- $examples ) where import Control.Concurrent (threadWaitWrite,threadWaitRead)-import Control.Exception (mask,onException)+import Control.Exception (throwIO,mask,onException) import Data.Primitive (ByteArray,MutableByteArray(..)) import Data.Word (Word16)-import Foreign.C.Error (Errno(..),eWOULDBLOCK,eAGAIN)+import Foreign.C.Error (Errno(..),eWOULDBLOCK,eAGAIN,eACCES) import Foreign.C.Types (CInt,CSize)-import GHC.Exts (Int(I#),RealWorld,shrinkMutableByteArray#)+import GHC.Exts (Int(I#),RealWorld,shrinkMutableByteArray#,ByteArray#,touch#)+import GHC.IO (IO(..)) import Net.Types (IPv4(..))-import Socket (SocketException(..),Context(..),Reason(..))+import Socket (SocketException(..),SocketUnrecoverableException(..),Direction(..),Interruptibility(..))+import Socket (cgetsockname)+import Socket.Datagram (SendException(..),ReceiveException(..))+import Socket.Datagram.IPv4.Undestined.Internal (Message(..),Socket(..))+import Socket.Datagram.IPv4.Undestined.Multiple (receiveMany,receiveManyUnless) import Socket.Debug (debug)-import Socket.IPv4 (Endpoint(..))-import System.Posix.Types (Fd)+import Socket.IPv4 (Endpoint(..),describeEndpoint) +import qualified Socket as SCK import qualified Control.Monad.Primitive as PM import qualified Data.Primitive as PM import qualified Linux.Socket as L import qualified Posix.Socket as S --- | A connectionless datagram socket that may communicate with many different--- endpoints on a datagram-by-datagram basis.-newtype Socket = Socket Fd- deriving (Eq,Ord)- -- | Open a socket and run the supplied callback on it. This closes the socket -- when the callback finishes or when an exception is thrown. Do not return -- the socket from the callback. This leads to undefined behavior. If the@@ -59,48 +64,66 @@ -> (Socket -> Word16 -> IO a) -- ^ Callback providing the socket and the chosen port -> IO (Either SocketException a) withSocket endpoint@Endpoint{port = specifiedPort} f = mask $ \restore -> do- debug ("withSocket: opening socket " ++ show endpoint)+ debug ("withSocket: opening socket " ++ describeEndpoint endpoint) e1 <- S.uninterruptibleSocket S.internet (L.applySocketFlags (L.closeOnExec <> L.nonblocking) S.datagram) S.defaultProtocol- debug ("withSocket: opened socket " ++ show endpoint)+ debug ("withSocket: opened socket " ++ describeEndpoint endpoint) case e1 of- Left err -> pure (Left (errorCode Open err))+ Left err -> throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Undestined+ functionWithSocket+ ["socket",describeEndpoint endpoint,describeErrorCode err] Right fd -> do e2 <- S.uninterruptibleBind fd (S.encodeSocketAddressInternet (endpointToSocketAddressInternet endpoint))- debug ("withSocket: requested binding for " ++ show endpoint)+ debug ("withSocket: requested binding for " ++ describeEndpoint endpoint) case e2 of Left err -> do -- We intentionally discard any exceptions thrown by close. There is -- simply nothing that can be done with them. S.uninterruptibleErrorlessClose fd- pure (Left (errorCode Bind err))+ throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Undestined+ functionWithSocket+ ["bind",describeEndpoint endpoint,describeErrorCode err] Right _ -> do eactualPort <- if specifiedPort == 0 then S.uninterruptibleGetSocketName fd S.sizeofSocketAddressInternet >>= \case Left err -> do S.uninterruptibleErrorlessClose fd- pure (Left (errorCode GetName err))+ throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Undestined+ functionWithSocket+ ["getsockname",describeEndpoint endpoint,describeErrorCode err] Right (sockAddrRequiredSz,sockAddr) -> if sockAddrRequiredSz == S.sizeofSocketAddressInternet then case S.decodeSocketAddressInternet sockAddr of Just S.SocketAddressInternet{port = actualPort} -> do let cleanPort = S.networkToHostShort actualPort- debug ("withSocket: successfully bound " ++ show endpoint ++ " and got port " ++ show cleanPort)+ debug ("withSocket: successfully bound " ++ describeEndpoint endpoint ++ " and got port " ++ show cleanPort) pure (Right cleanPort) Nothing -> do S.uninterruptibleErrorlessClose fd- pure (Left (exception GetName SocketAddressFamily))+ throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Undestined+ functionWithSocket+ [cgetsockname,describeEndpoint endpoint,"non-internet socket family"] else do S.uninterruptibleErrorlessClose fd- pure (Left (exception GetName SocketAddressSize))+ throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Undestined+ functionWithSocket+ [cgetsockname,describeEndpoint endpoint,"socket address size"] else pure (Right specifiedPort) case eactualPort of Left err -> pure (Left err) Right actualPort -> do a <- onException (restore (f (Socket fd) actualPort)) (S.uninterruptibleErrorlessClose fd) S.uninterruptibleClose fd >>= \case- Left err -> pure (Left (errorCode Close err))+ Left err -> throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Undestined+ functionWithSocket+ ["close",describeEndpoint endpoint,describeErrorCode err] Right _ -> pure (Right a) -- | Send a slice of a bytearray to the specified endpoint.@@ -111,44 +134,90 @@ -> Int -- ^ Offset into payload -> Int -- ^ Lenth of slice into buffer -> IO (Either SocketException ())-send (Socket !s) !remote !payload !off !len = do- debug ("send: about to send to " ++ show remote)- e1 <- S.uninterruptibleSendToByteArray s payload+send (Socket !s) !theRemote !thePayload !off !len = do+ debug ("send: about to send to " ++ show theRemote)+ e1 <- S.uninterruptibleSendToByteArray s thePayload (intToCInt off) (intToCSize len) mempty- (S.encodeSocketAddressInternet (endpointToSocketAddressInternet remote))- debug ("send: just sent to " ++ show remote)+ (S.encodeSocketAddressInternet (endpointToSocketAddressInternet theRemote))+ debug ("send: just sent to " ++ show theRemote) case e1 of Left err1 -> if err1 == eWOULDBLOCK || err1 == eAGAIN then do- debug ("send: waiting to for write ready to send to " ++ show remote)+ debug ("send: waiting to for write ready to send to " ++ show theRemote) threadWaitWrite s- e2 <- S.uninterruptibleSendToByteArray s payload+ e2 <- S.uninterruptibleSendToByteArray s thePayload (intToCInt off) (intToCSize len) mempty- (S.encodeSocketAddressInternet (endpointToSocketAddressInternet remote))+ (S.encodeSocketAddressInternet (endpointToSocketAddressInternet theRemote)) case e2 of Left err2 -> do debug ("send: encountered error after sending")- pure (Left (errorCode Send err2))+ throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Undestined+ functionSend+ [show theRemote,describeErrorCode err2] Right sz -> if csizeToInt sz == len then pure (Right ())- else pure (Left (exception Send (MessageTruncated (csizeToInt sz) len)))- else pure (Left (errorCode Send err1))+ else pure (Left (SentMessageTruncated (csizeToInt sz)))+ else throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Undestined+ functionSend+ [show theRemote,describeErrorCode err1] Right sz -> if csizeToInt sz == len then do debug ("send: success") pure (Right ())- else pure (Left (exception Send (MessageTruncated (csizeToInt sz) len)))+ else pure (Left (SentMessageTruncated (csizeToInt sz))) +-- | Send a slice of a bytearray to the specified endpoint.+sendMutableByteArraySlice ::+ Socket -- ^ Socket+ -> Endpoint -- ^ Remote IPv4 address and port+ -> MutableByteArray RealWorld -- ^ Buffer (will be sliced)+ -> Int -- ^ Offset into payload+ -> Int -- ^ Lenth of slice into buffer+ -> IO (Either (SendException 'Uninterruptible) ())+sendMutableByteArraySlice (Socket !s) !theRemote !thePayload !off !len = do+ debug ("send mutable: about to send to " ++ show theRemote)+ e1 <- S.uninterruptibleSendToMutableByteArray s thePayload+ (intToCInt off)+ (intToCSize len)+ mempty+ (S.encodeSocketAddressInternet (endpointToSocketAddressInternet theRemote))+ debug ("send mutable: just sent to " ++ show theRemote)+ case e1 of+ Left err1 -> if err1 == eWOULDBLOCK || err1 == eAGAIN+ then do+ debug ("send mutable: waiting to for write ready to send to " ++ show theRemote)+ threadWaitWrite s+ e2 <- S.uninterruptibleSendToMutableByteArray s thePayload+ (intToCInt off)+ (intToCSize len)+ mempty+ (S.encodeSocketAddressInternet (endpointToSocketAddressInternet theRemote))+ case e2 of+ Left err2 -> do+ debug ("send mutable: encountered error after sending")+ handleSendException functionSendMutableByteArray err2+ Right sz -> if csizeToInt sz == len+ then pure (Right ())+ else pure (Left (SendTruncated (csizeToInt sz)))+ else handleSendException functionSendMutableByteArray err1+ Right sz -> if csizeToInt sz == len+ then do+ debug ("send mutable: success")+ pure (Right ())+ else pure (Left (SendTruncated (csizeToInt sz)))+ -- | Receive a datagram into a freshly allocated bytearray.-receive ::+receiveByteArray :: Socket -- ^ Socket -> Int -- ^ Maximum size of datagram to receive- -> IO (Either SocketException (Endpoint,ByteArray))-receive (Socket !fd) !maxSz = do+ -> IO (Either (ReceiveException 'Uninterruptible) Message)+receiveByteArray (Socket !fd) !maxSz = do debug "receive: about to wait" threadWaitRead fd debug "receive: socket is now readable"@@ -161,20 +230,26 @@ (intToCSize maxSz) (L.truncate) S.sizeofSocketAddressInternet debug "receive: finished reading from socket" case e of- Left err -> pure (Left (errorCode Receive err))+ Left err -> throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Undestined+ functionReceive+ [describeErrorCode err] Right (sockAddrRequiredSz,sockAddr,recvSz) -> if csizeToInt recvSz <= maxSz then if sockAddrRequiredSz == S.sizeofSocketAddressInternet then case S.decodeSocketAddressInternet sockAddr of Just sockAddrInet -> do shrinkMutableByteArray marr (csizeToInt recvSz) arr <- PM.unsafeFreezeByteArray marr- pure $ Right- ( socketAddressInternetToEndpoint sockAddrInet- , arr- )- Nothing -> pure (Left (exception Receive SocketAddressFamily))- else pure (Left (exception Receive SocketAddressSize))- else pure (Left (exception Receive (MessageTruncated maxSz (csizeToInt recvSz))))+ pure $ Right (Message (socketAddressInternetToEndpoint sockAddrInet) arr)+ Nothing -> throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Undestined+ functionReceive+ [SCK.crecvfrom,SCK.nonInternetSocketFamily]+ else throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Undestined+ functionReceive+ [SCK.crecvfrom,SCK.socketAddressSize]+ else pure (Left (ReceiveTruncated (csizeToInt recvSz))) -- | Receive a datagram into a mutable byte array, ignoring information about -- the remote endpoint. Returns the actual number of bytes present in the@@ -193,10 +268,13 @@ -- exception. e <- S.uninterruptibleReceiveFromMutableByteArray_ fd buf (intToCInt off) (intToCSize maxSz) (L.truncate) case e of- Left err -> pure (Left (errorCode Receive err))+ Left err -> throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Undestined+ functionReceiveMutableByteArray+ [describeErrorCode err] Right recvSz -> if csizeToInt recvSz <= maxSz then pure (Right (csizeToInt recvSz))- else pure (Left (exception Receive (MessageTruncated maxSz (csizeToInt recvSz))))+ else pure (Left (ReceivedMessageTruncated (csizeToInt recvSz))) -- TODO: add receiveTimeout -- receiveTimeout ::@@ -218,20 +296,8 @@ , port = S.networkToHostShort port } -intToCInt :: Int -> CInt-intToCInt = fromIntegral--intToCSize :: Int -> CSize-intToCSize = fromIntegral--csizeToInt :: CSize -> Int-csizeToInt = fromIntegral--errorCode :: Context -> Errno -> SocketException-errorCode func (Errno x) = SocketException func (ErrorCode x)--exception :: Context -> Reason -> SocketException-exception func reason = SocketException func reason+errorCode :: Errno -> SocketException+errorCode (Errno x) = ErrorCode x shrinkMutableByteArray :: MutableByteArray RealWorld -> Int -> IO () shrinkMutableByteArray (MutableByteArray arr) (I# sz) =@@ -253,11 +319,58 @@ > unhandled $ withSocket (Endpoint IPv4.loopback 0) $ \sock port -> do > BC.putStrLn ("Receiving datagrams on 127.0.0.1:" <> BC.pack (show port)) > replicateM_ 10 $ do-> (remote,ByteArray payload) <- unhandled (receive sock 1024)-> BC.putStrLn ("Datagram from " <> BC.pack (show remote))-> BC.putStr (SB.fromShort (SB.SBS payload))+> DIU.Message sender (ByteArray contents) <- unhandled (DIU.receive sock 1024)+> BC.putStrLn ("Datagram from " <> BC.pack (show sender))+> BC.putStr (SB.fromShort (SB.SBS contents)) > > unhandled :: Exception e => IO (Either e a) -> IO a > unhandled action = action >>= either throwIO pure -}+++touchByteArray :: ByteArray -> IO ()+touchByteArray (PM.ByteArray x) = touchByteArray# x++touchByteArray# :: ByteArray# -> IO ()+touchByteArray# x = IO $ \s -> case touch# x s of s' -> (# s', () #)++intToCInt :: Int -> CInt+intToCInt = fromIntegral++intToCSize :: Int -> CSize+intToCSize = fromIntegral++csizeToInt :: CSize -> Int+csizeToInt = fromIntegral++moduleSocketDatagramIPv4Undestined :: String+moduleSocketDatagramIPv4Undestined = "Socket.Datagram.IPv4.Undestined"++functionReceive :: String+functionReceive = "receive"++functionSend :: String+functionSend = "send"++functionSendMutableByteArray :: String+functionSendMutableByteArray = "sendMutableByteArray"++functionReceiveMutableByteArray :: String+functionReceiveMutableByteArray = "receiveMutableByteArray"++functionWithSocket :: String+functionWithSocket = "withSocket"++describeErrorCode :: Errno -> String+describeErrorCode (Errno e) = "error code " ++ show e++handleSendException :: String -> Errno -> IO (Either (SendException i) a)+{-# INLINE handleSendException #-}+handleSendException func e+ | e == eACCES = pure (Left SendBroadcasted)+ | otherwise = throwIO $ SocketUnrecoverableException+ moduleSocketDatagramIPv4Undestined+ func+ [describeErrorCode e]+
+ src/Socket/Datagram/IPv4/Undestined/Internal.hs view
@@ -0,0 +1,24 @@+{-# language BangPatterns #-}+{-# language DeriveAnyClass #-}+{-# language DerivingStrategies #-}++module Socket.Datagram.IPv4.Undestined.Internal+ ( Socket(..)+ , Message(..)+ ) where++import Socket.IPv4 (Endpoint)+import System.Posix.Types (Fd)+import Data.Primitive (ByteArray)++-- | A connectionless datagram socket that may communicate with many different+-- endpoints on a datagram-by-datagram basis.+newtype Socket = Socket Fd+ deriving stock (Eq,Ord,Show)++data Message = Message+ { remote :: {-# UNPACK #-} !Endpoint+ , payload :: !ByteArray+ } deriving stock (Eq,Show)++
src/Socket/IPv4.hs view
@@ -1,14 +1,27 @@ {-# language BangPatterns #-}+{-# language DataKinds #-}+{-# language DeriveAnyClass #-} {-# language DerivingStrategies #-} {-# language DuplicateRecordFields #-}+{-# language GADTs #-}+{-# language KindSignatures #-}+{-# language NamedFieldPuns #-}+{-# language StandaloneDeriving #-} module Socket.IPv4 ( Endpoint(..)+ , SocketException(..)+ , describeEndpoint ) where +import Control.Exception (Exception)+import Data.Kind (Type) import Data.Word (Word16) import Net.Types (IPv4(..)) +import qualified Data.Text as T+import qualified Net.IPv4 as IPv4+ -- | An endpoint for an IPv4 socket, connection, or listener. -- Everything is in host byte order, and the user is not -- responisble for performing any conversions.@@ -16,3 +29,42 @@ { address :: !IPv4 , port :: !Word16 } deriving stock (Eq,Show)++-- This is used internally for debug messages and for presenting+-- unrecoverable exceptions.+describeEndpoint :: Endpoint -> String+describeEndpoint (Endpoint {address,port}) =+ T.unpack (IPv4.encode address) ++ ":" ++ show port++-- | Recoverable exceptions that happen when establishing an internet-domain+-- stream listener or datagram socket.+--+-- ==== __Discussion__+--+-- The recoverable exceptions that we from stream sockets (established+-- with @socket@-@bind@-@listen@) and datagram sockets (established with+-- @socket@-@bind@) are the exact same exceptions. Consequently, we reuse+-- the same type in both case. It is a little unfortunate since the name+-- @ListenException@ would be more appropriate for stream sockets. But+-- the code reuse is worth the naming quibble.+data SocketException :: Type where+ -- | The address is protected, and the user is not the superuser. This most+ -- commonly happens when trying to bind to a port below 1024. On Linux,+ -- When it is necessary to bind to such a port on Linux, consider using the+ -- <http://man7.org/linux/man-pages/man7/capabilities.7.html CAP_NET_BIND_SERVICE>+ -- capability instead of running the process as root. (@EACCES@)+ SocketPermissionDenied :: SocketException+ -- | The given address is already in use. (@EADDRINUSE@ with specified port)+ SocketAddressInUse :: SocketException+ -- | The port number was specified as zero, but upon attempting to+ -- bind to an ephemeral port, it was determined that all port numbers+ -- numbers in the ephemeral port range are currently in use.+ -- (@EADDRINUSE@ with unspecified port)+ SocketEphemeralPortsExhausted :: SocketException+ -- | A limit on the number of open file descriptors has been reached.+ -- This could be the per-process limit or the system limit.+ -- (@EMFILE@ and @ENFILE@)+ SocketFileDescriptorLimit :: SocketException++deriving stock instance Show SocketException+deriving anyclass instance Exception SocketException
+ src/Socket/Stream.hs view
@@ -0,0 +1,167 @@+{-# language DataKinds #-}+{-# language DeriveAnyClass #-}+{-# language DerivingStrategies #-}+{-# language GADTs #-}+{-# language KindSignatures #-}+{-# language StandaloneDeriving #-}+module Socket.Stream+ ( SendException(..)+ , ReceiveException(..)+ , ConnectException(..)+ , SocketException(..)+ , AcceptException(..)+ , CloseException(..)+ ) where++import Socket (Direction(..),Interruptibility(..),Forkedness(..))+import Socket.IPv4 (SocketException(..))++import Data.Kind (Type)+import Data.Typeable (Typeable)+import Control.Exception (Exception)++-- | Recoverable exceptions that can occur while connecting to a peer.+-- This includes both failures while opening the socket and failures+-- while connecting to the peer.+--+-- ==== __Discussion__+--+-- In its API for connecting to a peer, this library combines the step of+-- creating a socket with the step of connecting to the peer. In other words,+-- the end user never gets access to an unconnected stream socket.+-- Consequently, the connection exceptions correspond to the @socket@+-- errors @EMFILE@ and @ENFILE@ as well as the @connect@+-- errors @ECONNREFUSED@, @EACCES@/@EPERM@, @ETIMEDOUT@, @ENETUNREACH@, and+-- @EADDRNOTAVAIL@.+--+-- Somewhat surprisingly, @EADDRINUSE@ is not included in the list of @connect@+-- error codes we recognize as recoverable. The+-- <http://man7.org/linux/man-pages/man2/connect.2.html accept man page>+-- describes @EADDRINUSE@ as "Local address is already in use". However,+-- it is unclear what this means. The caller of @connect@ does not provide+-- an internet socket address. If ephemeral ports are exhausted, @connect@+-- will error with @EADDRNOTAVAIL@. An unresolved+-- <https://stackoverflow.com/questions/43199021/how-could-connect-fail-and-set-errno-to-eaddrinuse Stack Overflow question>+-- calls into question whether or not it is actually possible for this+-- error to happen with an internet domain socket. The author has decided+-- to omit any checks for it. This means that, if it does ever happen,+-- it will cause a @SocketUnrecoverableException@ to be thrown. The Linux+-- cognoscenti are encouraged to open an issue if they have more information+-- about the circumstances under which this exception can occur.+data ConnectException :: Interruptibility -> Type where+ -- | Either the connection was blocked by a local firewall rule or it+ -- was blocked because it was to a broadcast address. Sadly, these+ -- two errors are not distinguished by the Linux sockets API.+ -- (@EACCES@/@EPERM@)+ ConnectFirewalled :: ConnectException i+ -- | A limit on the number of open file descriptors has been reached.+ -- This could be the per-process limit or the system limit.+ -- (@EMFILE@ and @ENFILE@)+ ConnectFileDescriptorLimit :: ConnectException i+ -- | The network is unreachable. (@ENETUNREACH@)+ ConnectNetworkUnreachable :: ConnectException i+ -- | All port numbers numbers in the ephemeral port range are currently in+ -- use. (@EADDRNOTAVAIL@)+ ConnectEphemeralPortsExhausted :: ConnectException i+ -- | No one is listening on the remote address. (@ECONNREFUSED@)+ ConnectRefused :: ConnectException i+ -- | Timeout while attempting connection. The server may be too busy+ -- to accept new connections. Note that stock Linux configuration has+ -- timeout at + -- <http://willbryant.net/overriding_the_default_linux_kernel_20_second_tcp_socket_connect_timeout appropriately 20 seconds>.+ -- Users interested in timing out more quickly are encouraged to+ -- use @registerDelay@ with the @interruptible@ variants of the+ -- connection functions in this library. (@ETIMEDOUT@)+ ConnectTimeout :: ConnectException i+ -- | STM-style interrupt (much safer than C-style interrupt)+ ConnectInterrupted :: ConnectException 'Interruptible++deriving stock instance Show (ConnectException i)+deriving anyclass instance Typeable i => Exception (ConnectException i)+++data CloseException :: Type where+ -- | After the local process shut down the writing channel, it+ -- was expecting the peer to do the same. However, the peer+ -- sent more data instead. If this happens, the local process+ -- does still close the socket. However, it must send a TCP+ -- reset to accomplish this since there is still unread data+ -- in the receive buffer.+ --+ -- This can happen if the peer is misbehaving or if the consumer+ -- of the @sockets@ API has incorrectly implemented a protocol+ -- living above layer 4 of the OSI model.+ ClosePeerContinuedSending :: CloseException++deriving stock instance Show CloseException+deriving anyclass instance Exception CloseException++-- | Recoverable exceptions that can occur while accepting an inbound+-- connection.+data AcceptException :: Interruptibility -> Type where+ -- | The peer reset the connection before the running process+ -- accepted it. This is not typically treated as fatal. The+ -- process may continue accepting connections. (@ECONNABORTED@)+ AcceptConnectionAborted :: AcceptException i+ -- | A limit on the number of open file descriptors has been reached.+ -- This could be the per-process limit or the system limit.+ -- (@EMFILE@ and @ENFILE@)+ AcceptFileDescriptorLimit :: AcceptException i+ -- | Firewall rules forbid connection. (@EPERM@)+ AcceptFirewalled :: AcceptException i+ -- | STM-style interrupt (much safer than C-style interrupt)+ AcceptInterrupted :: AcceptException 'Interruptible++deriving stock instance Show (AcceptException i)+deriving anyclass instance (Typeable i) => Exception (AcceptException i)+++data SendException :: Interruptibility -> Type where+ -- | The local socket has already shutdown its writing channel.+ -- Consequently, sending is no longer possible. This can happen+ -- even if the process does not @shutdown@ the socket. If the+ -- peer decides to @close@ the connection, the local operating system+ -- will shutdown both the reading and writing channels. (@EPIPE@)+ SendShutdown :: SendException i+ -- | The peer reset the connection.+ SendReset :: SendException i+ -- | STM-style interrupt (much safer than C-style interrupt)+ SendInterrupted :: SendException 'Interruptible++deriving stock instance Show (SendException i)+deriving anyclass instance Typeable i => Exception (SendException i)++-- | Recoverable exceptions that can occur while receiving data on a+-- stream socket.+--+-- ==== __Discussion__+--+-- The <http://man7.org/linux/man-pages/man2/recv.2.html recv man page>+-- explicitly documents these:+--+-- * @EAGAIN@/@EAGAIN@: Not possible after using event manager to wait.+-- * @EBADF@: Prevented by this library.+-- * @ECONNREFUSED@: Not sure if this is possible. Currently treated as+-- an unrecoverable exception.+-- * @EFAULT@: Not recoverable. API consumer has misused @Addr@.+-- * @EINTR@: Prevented by this library. Unsafe FFI is not interruptible.+-- * @EINVAL@: Prevented by this library.+-- * @ENOMEM@: Not recoverable.+-- * @ENOTCONN@: Prevented by this library.+-- * @ENOTSOCK@: Prevented by this library.+--+-- The man page includes a disclaimer: "Additional errors may be generated+-- and returned from the underlying protocol modules". One such error+-- when dealing with stream sockets in @ECONNRESET@. One scenario where+-- this happens is when the process running on the peer terminates ungracefully+-- and the operating system on the peer cleans up by sending a reset.+data ReceiveException :: Interruptibility -> Type where+ -- | The peer shutdown its writing channel. (zero-length chunk)+ ReceiveShutdown :: ReceiveException i+ -- | The peer reset the connection. (@ECONNRESET@)+ ReceiveReset :: ReceiveException i+ -- | STM-style interrupt (much safer than C-style interrupt)+ ReceiveInterrupted :: ReceiveException 'Interruptible++deriving stock instance Show (ReceiveException i)+deriving anyclass instance Typeable i => Exception (ReceiveException i)
src/Socket/Stream/IPv4.hs view
@@ -1,493 +1,1211 @@-{-# language BangPatterns #-}-{-# language RankNTypes #-}-{-# language DuplicateRecordFields #-}-{-# language LambdaCase #-}-{-# language NamedFieldPuns #-}-{-# language MagicHash #-}--module Socket.Stream.IPv4- ( -- * Types- Listener- , Connection- , Endpoint(..)- -- * Bracketed- , withListener- , withAccepted- , withConnection- , forkAccepted- , forkAcceptedUnmasked- -- * Communicate- , sendByteArray- , sendByteArraySlice- , sendMutableByteArray- , sendMutableByteArraySlice- , receiveByteArray- , receiveBoundedByteArray- , receiveMutableByteArray- -- * Exceptions- , SocketException(..)- , Context(..)- , Reason(..)- ) where--import Control.Concurrent (ThreadId,threadWaitWrite,threadWaitRead)-import Control.Concurrent (forkIO,forkIOWithUnmask)-import Control.Exception (mask,onException)-import Data.Bifunctor (bimap)-import Data.Primitive (ByteArray,MutableByteArray(..))-import Data.Word (Word16)-import Foreign.C.Error (Errno(..),eAGAIN,eWOULDBLOCK,eINPROGRESS)-import Foreign.C.Types (CInt,CSize)-import GHC.Exts (RealWorld,Int(I#),shrinkMutableByteArray#)-import Socket (SocketException(..),Context(..),Reason(..))-import Socket.Debug (debug)-import Socket.IPv4 (Endpoint(..))-import System.Posix.Types (Fd)-import Net.Types (IPv4(..))--import qualified Control.Monad.Primitive as PM-import qualified Data.Primitive as PM-import qualified Linux.Socket as L-import qualified Posix.Socket as S---- | A socket that listens for incomming connections.-newtype Listener = Listener Fd---- | A connection-oriented stream socket.-newtype Connection = Connection Fd--withListener ::- Endpoint- -> (Listener -> Word16 -> IO a)- -> IO (Either SocketException a)-withListener endpoint@Endpoint{port = specifiedPort} f = mask $ \restore -> do- debug ("withSocket: opening listener " ++ show endpoint)- e1 <- S.uninterruptibleSocket S.internet- (L.applySocketFlags (L.closeOnExec <> L.nonblocking) S.stream)- S.defaultProtocol- debug ("withSocket: opened listener " ++ show endpoint)- case e1 of- Left err -> pure (Left (errorCode Open err))- Right fd -> do- e2 <- S.uninterruptibleBind fd- (S.encodeSocketAddressInternet (endpointToSocketAddressInternet endpoint))- debug ("withSocket: requested binding for listener " ++ show endpoint)- case e2 of- Left err -> do- _ <- S.uninterruptibleClose fd- pure (Left (errorCode Bind err))- Right _ -> S.uninterruptibleListen fd 16 >>= \case- -- We hardcode the listen backlog to 16. The author is unfamiliar- -- with use cases where gains are realized from tuning this parameter.- -- Open an issue if this causes problems for anyone.- Left err -> do- _ <- S.uninterruptibleClose fd- debug "withSocket: listen failed with error code"- pure (Left (errorCode Listen err))- Right _ -> do - -- The getsockname is copied from code in Socket.Datagram.IPv4.Undestined.- -- Consider factoring this out.- eactualPort <- if specifiedPort == 0- then S.uninterruptibleGetSocketName fd S.sizeofSocketAddressInternet >>= \case- Left err -> do- _ <- S.uninterruptibleClose fd- pure (Left (errorCode GetName err))- Right (sockAddrRequiredSz,sockAddr) -> if sockAddrRequiredSz == S.sizeofSocketAddressInternet- then case S.decodeSocketAddressInternet sockAddr of- Just S.SocketAddressInternet{port = actualPort} -> do- let cleanActualPort = S.networkToHostShort actualPort- debug ("withSocket: successfully bound listener " ++ show endpoint ++ " and got port " ++ show cleanActualPort)- pure (Right cleanActualPort)- Nothing -> do- _ <- S.uninterruptibleClose fd- pure (Left (exception GetName SocketAddressFamily))- else do- _ <- S.uninterruptibleClose fd- pure (Left (exception GetName SocketAddressSize))- else pure (Right specifiedPort)- case eactualPort of- Left err -> pure (Left err)- Right actualPort -> do- a <- onException (restore (f (Listener fd) actualPort)) (S.uninterruptibleClose fd)- S.uninterruptibleClose fd >>= \case- Left err -> pure (Left (errorCode Close err))- Right _ -> pure (Right a)---- | Accept a connection on the listener and run the supplied callback--- on it. This closes the connection when the callback finishes or if--- an exception is thrown. Since this function blocks the thread until--- the callback finishes, it is only suitable for stream socket clients--- that handle one connection at a time. The variant 'forkAcceptedUnmasked'--- is preferrable for servers that need to handle connections concurrently--- (most use cases).-withAccepted ::- Listener- -> (Connection -> Endpoint -> IO a)- -> IO (Either SocketException a)-withAccepted lst cb = internalAccepted- ( \restore action -> do- action restore- ) lst cb--internalAccepted ::- ((forall x. IO x -> IO x) -> ((IO a -> IO b) -> IO (Either SocketException b)) -> IO (Either SocketException c))- -> Listener- -> (Connection -> Endpoint -> IO a)- -> IO (Either SocketException c)-internalAccepted wrap (Listener !lst) f = do- threadWaitRead lst- mask $ \restore -> do- S.uninterruptibleAccept lst S.sizeofSocketAddressInternet >>= \case- Left err -> pure (Left (errorCode Accept err))- Right (sockAddrRequiredSz,sockAddr,acpt) -> if sockAddrRequiredSz == S.sizeofSocketAddressInternet- then case S.decodeSocketAddressInternet sockAddr of- Just sockAddrInet -> do- let acceptedEndpoint = socketAddressInternetToEndpoint sockAddrInet- debug ("withAccepted: successfully accepted connection from " ++ show acceptedEndpoint)- wrap restore $ \restore' -> do- a <- onException (restore' (f (Connection acpt) acceptedEndpoint)) (S.uninterruptibleClose acpt)- gracefulClose acpt a- Nothing -> do- _ <- S.uninterruptibleClose acpt- pure (Left (exception GetName SocketAddressFamily))- else do- _ <- S.uninterruptibleClose acpt- pure (Left (exception GetName SocketAddressSize))--gracefulClose :: Fd -> a -> IO (Either SocketException a)-gracefulClose fd a = S.uninterruptibleShutdown fd S.write >>= \case- Left err -> do- _ <- S.uninterruptibleClose fd- pure (Left (errorCode Shutdown err))- Right _ -> do- buf <- PM.newByteArray 1- S.uninterruptibleReceiveMutableByteArray fd buf 0 1 mempty >>= \case- Left err1 -> if err1 == eWOULDBLOCK || err1 == eAGAIN- then do- threadWaitRead fd- S.uninterruptibleReceiveMutableByteArray fd buf 0 1 mempty >>= \case- Left err -> do- _ <- S.uninterruptibleClose fd- pure (Left (errorCode Shutdown err))- Right sz -> if sz == 0- then fmap (bimap (errorCode Close) (const a)) (S.uninterruptibleClose fd)- else do- debug ("Socket.Stream.IPv4.gracefulClose: remote not shutdown A")- _ <- S.uninterruptibleClose fd- pure (Left (exception Shutdown RemoteNotShutdown))- else do- _ <- S.uninterruptibleClose fd- -- Is this the right error context? It's a call- -- to recv, but it happens while shutting down- -- the socket.- pure (Left (errorCode Shutdown err1))- Right sz -> if sz == 0- then fmap (bimap (errorCode Close) (const a)) (S.uninterruptibleClose fd)- else do- debug ("Socket.Stream.IPv4.gracefulClose: remote not shutdown B")- _ <- S.uninterruptibleClose fd- pure (Left (exception Shutdown RemoteNotShutdown))---- | Accept a connection on the listener and run the supplied callback in--- a new thread. Prefer 'forkAcceptedUnmasked' unless the masking state--- needs to be preserved for the callback. Such a situation seems unlikely--- to the author.-forkAccepted ::- Listener- -> (Either SocketException a -> IO ())- -> (Connection -> Endpoint -> IO a)- -> IO (Either SocketException ThreadId)-forkAccepted lst consumeException cb = internalAccepted- ( \restore action -> do- tid <- forkIO $ do- x <- action restore- restore (consumeException x)- pure (Right tid)- ) lst cb---- | Accept a connection on the listener and run the supplied callback in--- a new thread. The masking state is set to @Unmasked@ when running the--- callback.-forkAcceptedUnmasked ::- Listener- -> (Either SocketException a -> IO ())- -> (Connection -> Endpoint -> IO a)- -> IO (Either SocketException ThreadId)-forkAcceptedUnmasked lst consumeException cb = internalAccepted- ( \_ action -> do- tid <- forkIOWithUnmask $ \unmask -> do- x <- action unmask- unmask (consumeException x)- pure (Right tid)- ) lst cb---- | Establish a connection to a server.-withConnection ::- Endpoint -- ^ Remote endpoint- -> (Connection -> IO a) -- ^ Callback to consume connection- -> IO (Either SocketException a)-withConnection !remote f = mask $ \restore -> do- debug ("withSocket: opening connection " ++ show remote)- e1 <- S.uninterruptibleSocket S.internet- (L.applySocketFlags (L.closeOnExec <> L.nonblocking) S.stream)- S.defaultProtocol- debug ("withSocket: opened connection " ++ show remote)- case e1 of- Left err1 -> pure (Left (errorCode Open err1))- Right fd -> do- let sockAddr = id- $ S.encodeSocketAddressInternet- $ endpointToSocketAddressInternet- $ remote- merr <- S.uninterruptibleConnect fd sockAddr >>= \case- Left err2 -> if err2 == eINPROGRESS- then do- threadWaitWrite fd- pure Nothing- else pure (Just (errorCode Connect err2))- Right _ -> pure Nothing- case merr of- Just err -> do- _ <- S.uninterruptibleClose fd- pure (Left err)- Nothing -> do- e <- S.uninterruptibleGetSocketOption fd- S.levelSocket S.optionError (intToCInt (PM.sizeOf (undefined :: CInt)))- case e of- Left err -> do- _ <- S.uninterruptibleClose fd- pure (Left (errorCode Option err))- Right (sz,S.OptionValue val) -> if sz == intToCInt (PM.sizeOf (undefined :: CInt))- then- let err = PM.indexByteArray val 0 :: CInt in- if err == 0- then do- a <- onException (restore (f (Connection fd))) (S.uninterruptibleClose fd)- gracefulClose fd a- else do- _ <- S.uninterruptibleClose fd- pure (Left (errorCode Connect (Errno err)))- else do- _ <- S.uninterruptibleClose fd- pure (Left (exception Option OptionValueSize))--sendByteArray ::- Connection -- ^ Connection- -> ByteArray -- ^ Buffer (will be sliced)- -> IO (Either SocketException ())-sendByteArray conn arr =- sendByteArraySlice conn arr 0 (PM.sizeofByteArray arr)--sendByteArraySlice ::- Connection -- ^ Connection- -> ByteArray -- ^ Buffer (will be sliced)- -> Int -- ^ Offset into payload- -> Int -- ^ Lenth of slice into buffer- -> IO (Either SocketException ())-sendByteArraySlice !conn !payload !off0 !len0 = go off0 len0- where- go !off !len = if len > 0- then internalSend conn payload off len >>= \case- Left e -> pure (Left e)- Right sz' -> do- let sz = csizeToInt sz'- go (off + sz) (len - sz)- else pure (Right ())--sendMutableByteArray ::- Connection -- ^ Connection- -> MutableByteArray RealWorld -- ^ Buffer (will be sliced)- -> IO (Either SocketException ())-sendMutableByteArray conn arr =- sendMutableByteArraySlice conn arr 0 =<< PM.getSizeofMutableByteArray arr--sendMutableByteArraySlice ::- Connection -- ^ Connection- -> MutableByteArray RealWorld -- ^ Buffer (will be sliced)- -> Int -- ^ Offset into payload- -> Int -- ^ Lenth of slice into buffer- -> IO (Either SocketException ())-sendMutableByteArraySlice !conn !payload !off0 !len0 = go off0 len0- where- go !off !len = if len > 0- then internalSendMutable conn payload off len >>= \case- Left e -> pure (Left e)- Right sz' -> do- let sz = csizeToInt sz'- go (off + sz) (len - sz)- else pure (Right ())---- The length must be greater than zero.-internalSendMutable :: - Connection -- ^ Connection- -> MutableByteArray RealWorld -- ^ Buffer (will be sliced)- -> Int -- ^ Offset into payload- -> Int -- ^ Length of slice into buffer- -> IO (Either SocketException CSize)-internalSendMutable (Connection !s) !payload !off !len = do- e1 <- S.uninterruptibleSendMutableByteArray s payload- (intToCInt off)- (intToCSize len)- mempty- case e1 of- Left err1 -> if err1 == eWOULDBLOCK || err1 == eAGAIN- then do- threadWaitWrite s- e2 <- S.uninterruptibleSendMutableByteArray s payload- (intToCInt off)- (intToCSize len)- mempty- case e2 of- Left err2 -> pure (Left (errorCode Send err2))- Right sz -> pure (Right sz)- else pure (Left (errorCode Send err1))- Right sz -> pure (Right sz)---- The length must be greater than zero.-internalSend ::- Connection -- ^ Connection- -> ByteArray -- ^ Buffer (will be sliced)- -> Int -- ^ Offset into payload- -> Int -- ^ Length of slice into buffer- -> IO (Either SocketException CSize)-internalSend (Connection !s) !payload !off !len = do- debug ("send: about to send chunk on stream socket, offset " ++ show off ++ " and length " ++ show len)- e1 <- S.uninterruptibleSendByteArray s payload- (intToCInt off)- (intToCSize len)- mempty- debug "send: just sent chunk on stream socket"- case e1 of- Left err1 -> if err1 == eWOULDBLOCK || err1 == eAGAIN- then do- debug "send: waiting to for write ready on stream socket"- threadWaitWrite s- e2 <- S.uninterruptibleSendByteArray s payload- (intToCInt off)- (intToCSize len)- mempty- case e2 of- Left err2 -> do- debug "send: encountered error after sending chunk on stream socket"- pure (Left (errorCode Send err2))- Right sz -> pure (Right sz)- else pure (Left (errorCode Send err1))- Right sz -> pure (Right sz)---- The maximum number of bytes to receive must be greater than zero.--- The operating system guarantees us that the returned actual number--- of bytes is less than or equal to the requested number of bytes.--- This function does not validate that the result size is greater--- than zero. Functions calling this must perform that check. This--- also does not trim the buffer. The caller must do that if it is--- necessary.-internalReceiveMaximally ::- Connection -- ^ Connection- -> Int -- ^ Maximum number of bytes to receive- -> MutableByteArray RealWorld -- ^ Receive buffer- -> Int -- ^ Offset into buffer- -> IO (Either SocketException Int)-internalReceiveMaximally (Connection !fd) !maxSz !buf !off = do- debug "receive: stream socket about to wait"- threadWaitRead fd- debug ("receive: stream socket is now readable, receiving up to " ++ show maxSz ++ " bytes at offset " ++ show off)- e <- S.uninterruptibleReceiveMutableByteArray fd buf (intToCInt off) (intToCSize maxSz) mempty- debug "receive: finished reading from stream socket"- case e of- Left err -> pure (Left (errorCode Receive err))- Right recvSz -> pure (Right (csizeToInt recvSz))---- | Receive exactly the given number of bytes. If the remote application--- shuts down its end of the connection before sending the required--- number of bytes, this returns--- @'Left' ('SocketException' 'Receive' 'RemoteShutdown')@.-receiveByteArray ::- Connection -- ^ Connection- -> Int -- ^ Number of bytes to receive- -> IO (Either SocketException ByteArray)-receiveByteArray !conn0 !total = do- marr <- PM.newByteArray total- go conn0 marr 0 total- where- go !conn !marr !off !remaining = case compare remaining 0 of- GT -> internalReceiveMaximally conn remaining marr off >>= \case- Left err -> pure (Left err)- Right sz -> if sz /= 0- then go conn marr (off + sz) (remaining - sz)- else pure (Left (exception Receive RemoteShutdown))- EQ -> do- arr <- PM.unsafeFreezeByteArray marr- pure (Right arr)- LT -> pure (Left (exception Receive NegativeBytesRequested))---- | Receive a number of bytes exactly equal to the size of the mutable--- byte array. If the remote application shuts down its end of the--- connection before sending the required number of bytes, this returns--- @'Left' ('SocketException' 'Receive' 'RemoteShutdown')@.-receiveMutableByteArray ::- Connection- -> MutableByteArray RealWorld- -> IO (Either SocketException ())-receiveMutableByteArray !conn0 !marr0 = do- total <- PM.getSizeofMutableByteArray marr0- go conn0 marr0 0 total- where- go !conn !marr !off !remaining = if remaining > 0- then internalReceiveMaximally conn remaining marr off >>= \case- Left err -> pure (Left err)- Right sz -> if sz /= 0- then go conn marr (off + sz) (remaining - sz)- else pure (Left (exception Receive RemoteShutdown))- else pure (Right ())---- | Receive up to the given number of bytes. If the remote application--- shuts down its end of the connection instead of sending any bytes,--- this returns--- @'Left' ('SocketException' 'Receive' 'RemoteShutdown')@.-receiveBoundedByteArray :: - Connection -- ^ Connection- -> Int -- ^ Maximum number of bytes to receive- -> IO (Either SocketException ByteArray)-receiveBoundedByteArray !conn !total- | total > 0 = do- m <- PM.newByteArray total- internalReceiveMaximally conn total m 0 >>= \case- Left err -> pure (Left err) - Right sz -> if sz /= 0- then do- shrinkMutableByteArray m sz- fmap Right (PM.unsafeFreezeByteArray m)- else pure (Left (exception Receive RemoteShutdown))- | total == 0 = pure (Right mempty)- | otherwise = pure (Left (exception Receive NegativeBytesRequested))--endpointToSocketAddressInternet :: Endpoint -> S.SocketAddressInternet-endpointToSocketAddressInternet (Endpoint {address, port}) = S.SocketAddressInternet- { port = S.hostToNetworkShort port- , address = S.hostToNetworkLong (getIPv4 address)- }--socketAddressInternetToEndpoint :: S.SocketAddressInternet -> Endpoint-socketAddressInternetToEndpoint (S.SocketAddressInternet {address,port}) = Endpoint- { address = IPv4 (S.networkToHostLong address)- , port = S.networkToHostShort port- }--errorCode :: Context -> Errno -> SocketException-errorCode func (Errno x) = SocketException func (ErrorCode x)--exception :: Context -> Reason -> SocketException-exception func reason = SocketException func reason--intToCInt :: Int -> CInt-intToCInt = fromIntegral--intToCSize :: Int -> CSize-intToCSize = fromIntegral--csizeToInt :: CSize -> Int-csizeToInt = fromIntegral--shrinkMutableByteArray :: MutableByteArray RealWorld -> Int -> IO ()-shrinkMutableByteArray (MutableByteArray arr) (I# sz) =- PM.primitive_ (shrinkMutableByteArray# arr sz)+{-# LANGUAGE BangPatterns #-}+{-# LANGUAGE DataKinds #-}+{-# LANGUAGE DuplicateRecordFields #-}+{-# LANGUAGE LambdaCase #-}+{-# LANGUAGE MagicHash #-}+{-# LANGUAGE NamedFieldPuns #-}+{-# LANGUAGE RankNTypes #-}+{-# LANGUAGE GADTs #-}++module Socket.Stream.IPv4+ ( -- * Types+ Listener+ , Connection+ , Endpoint(..)+ -- * Bracketed+ , withListener+ , withAccepted+ , withConnection+ , forkAccepted+ , forkAcceptedUnmasked+ , interruptibleForkAcceptedUnmasked+ -- * Communicate+ , sendByteArray+ , sendByteArraySlice+ , sendMutableByteArray+ , sendMutableByteArraySlice+ , interruptibleSendByteArray+ , interruptibleSendByteArraySlice+ , interruptibleSendMutableByteArraySlice+ , receiveByteArray+ , receiveBoundedByteArray+ , receiveBoundedMutableByteArraySlice+ , receiveMutableByteArray+ , interruptibleReceiveByteArray+ , interruptibleReceiveBoundedMutableByteArraySlice+ -- * Exceptions+ , SendException(..)+ , ReceiveException(..)+ , ConnectException(..)+ , SocketException(..)+ , AcceptException(..)+ , CloseException(..)+ , Interruptibility(..)+ -- * Unbracketed+ -- $unbracketed+ , listen+ , unlisten+ , unlisten_+ , connect+ , disconnect+ , disconnect_+ , accept+ , interruptibleAccept+ ) where++import Control.Applicative ((<|>))+import Control.Concurrent (ThreadId, threadWaitRead, threadWaitWrite)+import Control.Concurrent (threadWaitReadSTM,threadWaitWriteSTM)+import Control.Concurrent (forkIO, forkIOWithUnmask)+import Control.Exception (mask, mask_, onException, throwIO)+import Control.Monad.STM (STM,atomically,retry)+import Control.Concurrent.STM (TVar,modifyTVar',readTVar)+import Data.Bifunctor (bimap,first)+import Data.Bool (bool)+import Data.Functor (($>))+import Data.Primitive (ByteArray, MutableByteArray(..))+import Data.Word (Word16)+import Foreign.C.Error (Errno(..), eAGAIN, eINPROGRESS, eWOULDBLOCK, ePIPE, eNOTCONN)+import Foreign.C.Error (eADDRINUSE,eCONNRESET)+import Foreign.C.Error (eNFILE,eMFILE,eACCES,ePERM,eCONNABORTED)+import Foreign.C.Error (eTIMEDOUT,eADDRNOTAVAIL,eNETUNREACH,eCONNREFUSED)+import Foreign.C.Types (CInt, CSize)+import GHC.Exts (Int(I#), RealWorld, shrinkMutableByteArray#)+import Net.Types (IPv4(..))+import Socket (Interruptibility(..))+import Socket (SocketUnrecoverableException(..))+import Socket (cgetsockname,cclose)+import Socket.Debug (debug)+import Socket.IPv4 (Endpoint(..),describeEndpoint)+import Socket.Stream (ConnectException(..),SocketException(..),AcceptException(..))+import Socket.Stream (SendException(..),ReceiveException(..),CloseException(..))+import System.Posix.Types(Fd)++import qualified Control.Monad.Primitive as PM+import qualified Data.Primitive as PM+import qualified Linux.Socket as L+import qualified Posix.Socket as S+import qualified Socket as SCK++-- | A socket that listens for incomming connections.+newtype Listener = Listener Fd++-- | A connection-oriented stream socket.+newtype Connection = Connection Fd++-- | Open a socket that can be used to listen for inbound connections.+-- Requirements:+--+-- * This function may only be called in contexts where exceptions+-- are masked.+-- * The caller /must/ be sure to call 'unlistener' on the resulting+-- 'Listener' exactly once to close underlying file descriptor.+-- * The 'Listener' cannot be used after being given as an argument+-- to 'unlistener'.+--+-- Noncompliant use of this function leads to undefined behavior. Prefer+-- 'withListener' unless you are writing an integration with a+-- resource-management library.+listen :: Endpoint -> IO (Either SocketException (Listener, Word16))+listen endpoint@Endpoint{port = specifiedPort} = do+ debug ("listen: opening listen " ++ describeEndpoint endpoint)+ e1 <- S.uninterruptibleSocket S.internet+ (L.applySocketFlags (L.closeOnExec <> L.nonblocking) S.stream)+ S.defaultProtocol+ debug ("listen: opened listen " ++ describeEndpoint endpoint)+ case e1 of+ Left err -> handleSocketListenException SCK.functionWithListener err+ Right fd -> do+ e2 <- S.uninterruptibleBind fd+ (S.encodeSocketAddressInternet (endpointToSocketAddressInternet endpoint))+ debug ("listen: requested binding for listen " ++ describeEndpoint endpoint)+ case e2 of+ Left err -> do+ _ <- S.uninterruptibleClose fd+ handleBindListenException specifiedPort SCK.functionWithListener err+ Right _ -> S.uninterruptibleListen fd 16 >>= \case+ -- We hardcode the listen backlog to 16. The author is unfamiliar+ -- with use cases where gains are realized from tuning this parameter.+ -- Open an issue if this causes problems for anyone.+ Left err -> do+ _ <- S.uninterruptibleClose fd+ debug "listen: listen failed with error code"+ handleBindListenException specifiedPort SCK.functionWithListener err+ Right _ -> do+ -- The getsockname is copied from code in Socket.Datagram.IPv4.Undestined.+ -- Consider factoring this out.+ actualPort <- if specifiedPort == 0+ then S.uninterruptibleGetSocketName fd S.sizeofSocketAddressInternet >>= \case+ Left err -> throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ functionWithListener+ [cgetsockname,describeEndpoint endpoint,describeErrorCode err]+ Right (sockAddrRequiredSz,sockAddr) -> if sockAddrRequiredSz == S.sizeofSocketAddressInternet+ then case S.decodeSocketAddressInternet sockAddr of+ Just S.SocketAddressInternet{port = actualPort} -> do+ let cleanActualPort = S.networkToHostShort actualPort+ debug ("listen: successfully bound listen " ++ describeEndpoint endpoint ++ " and got port " ++ show cleanActualPort)+ pure cleanActualPort+ Nothing -> do+ _ <- S.uninterruptibleClose fd+ throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ functionWithListener+ [cgetsockname,"non-internet socket family"]+ else do+ _ <- S.uninterruptibleClose fd+ throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ functionWithListener+ [cgetsockname,describeEndpoint endpoint,"socket address size"]+ else pure specifiedPort+ pure (Right (Listener fd, actualPort))++-- | Close a listener. This throws an unrecoverable exception if+-- the socket cannot be closed.+unlisten :: Listener -> IO ()+unlisten (Listener fd) = S.uninterruptibleClose fd >>= \case+ Left err -> throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ functionWithListener+ [cclose,describeErrorCode err]+ Right _ -> pure ()++-- | Close a listener. This does not check to see whether or not+-- the operating system successfully closed the socket. It never+-- throws exceptions of any kind. This should only be preferred+-- to 'unlistener' in exception-cleanup contexts where there is+-- already an exception that will be rethrown. See the implementation+-- of 'withListener' for an example of appropriate use of both+-- 'unlistener' and 'unlistener_'.+unlisten_ :: Listener -> IO ()+unlisten_ (Listener fd) = S.uninterruptibleErrorlessClose fd++withListener ::+ Endpoint+ -> (Listener -> Word16 -> IO a)+ -> IO (Either SocketException a)+withListener endpoint f = mask $ \restore -> do+ listen endpoint >>= \case+ Left err -> pure (Left err)+ Right (sck, actualPort) -> do+ a <- onException+ (restore (f sck actualPort))+ (unlisten_ sck)+ unlisten sck+ pure (Right a)++-- | Listen for an inbound connection.+accept :: Listener -> IO (Either (AcceptException 'Uninterruptible) (Connection,Endpoint))+accept (Listener fd) = do+ -- Although this function must be called in a context where+ -- exceptions are masked, recall that threadWaitRead uses+ -- takeMVar, meaning that this first part is still interruptible.+ -- This is a good thing in the case of this function.+ threadWaitRead fd+ waitlessAccept fd++-- | Listen for an inbound connection. Can be interrupted by an+-- STM-style interrupt.+interruptibleAccept ::+ TVar Bool+ -- ^ Interrupted. If this becomes 'True' give up and return+ -- @'Left' 'AcceptInterrupted'@.+ -> Listener+ -> IO (Either (AcceptException 'Interruptible) (Connection,Endpoint))+interruptibleAccept abandon (Listener fd) = do+ interruptibleWaitRead abandon fd >>= \case+ True -> waitlessAccept fd+ False -> pure (Left AcceptInterrupted)++-- Only used internally+interruptibleAcceptCounting :: + TVar Int+ -> TVar Bool+ -> Listener+ -> IO (Either (AcceptException 'Interruptible) (Connection,Endpoint))+interruptibleAcceptCounting counter abandon (Listener fd) = do+ interruptibleWaitReadCounting counter abandon fd >>= \case+ True -> waitlessAccept fd+ False -> pure (Left AcceptInterrupted)++waitlessAccept :: Fd -> IO (Either (AcceptException i) (Connection,Endpoint))+waitlessAccept lstn = do+ S.uninterruptibleAccept lstn S.sizeofSocketAddressInternet >>= \case+ Left err -> handleAcceptException "withAccepted" err+ Right (sockAddrRequiredSz,sockAddr,acpt) -> if sockAddrRequiredSz == S.sizeofSocketAddressInternet+ then case S.decodeSocketAddressInternet sockAddr of+ Just sockAddrInet -> do+ let !acceptedEndpoint = socketAddressInternetToEndpoint sockAddrInet+ debug ("internalAccepted: successfully accepted connection from " ++ show acceptedEndpoint)+ pure (Right (Connection acpt, acceptedEndpoint))+ Nothing -> do+ _ <- S.uninterruptibleClose acpt+ throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ SCK.functionWithAccepted+ [SCK.cgetsockname,SCK.nonInternetSocketFamily]+ else do+ _ <- S.uninterruptibleClose acpt+ throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ SCK.functionWithAccepted+ [SCK.cgetsockname,SCK.socketAddressSize]++-- This function factors out the common elements of withAccepted, forkAccepted,+-- and forkAcceptedUnmasked. Unfortunately, I can barely understand it. The+-- higher-rank callback is particularly impenetrable. Sorry.+internalAccepted ::+ ((forall x. IO x -> IO x) -> ((IO a -> IO d) -> IO (Either CloseException (),d)) -> IO (Either (AcceptException 'Uninterruptible) c))+ -> Listener+ -> (Connection -> Endpoint -> IO a)+ -> IO (Either (AcceptException 'Uninterruptible) c)+internalAccepted wrap (Listener !lst) f = do+ threadWaitRead lst+ mask $ \restore -> do+ S.uninterruptibleAccept lst S.sizeofSocketAddressInternet >>= \case+ Left err -> handleAcceptException "withAccepted" err+ Right (sockAddrRequiredSz,sockAddr,acpt) -> if sockAddrRequiredSz == S.sizeofSocketAddressInternet+ then case S.decodeSocketAddressInternet sockAddr of+ Just sockAddrInet -> do+ let acceptedEndpoint = socketAddressInternetToEndpoint sockAddrInet+ debug ("internalAccepted: successfully accepted connection from " ++ show acceptedEndpoint)+ wrap restore $ \restore' -> do+ a <- onException (restore' (f (Connection acpt) acceptedEndpoint)) (S.uninterruptibleClose acpt)+ e <- gracefulCloseA acpt+ pure (e,a)+ Nothing -> do+ _ <- S.uninterruptibleClose acpt+ throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ SCK.functionWithAccepted+ [SCK.cgetsockname,SCK.nonInternetSocketFamily]+ else do+ _ <- S.uninterruptibleClose acpt+ throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ SCK.functionWithAccepted+ [SCK.cgetsockname,SCK.socketAddressSize]++gracefulCloseA :: Fd -> IO (Either CloseException ())+gracefulCloseA fd = S.uninterruptibleShutdown fd S.write >>= \case+ -- On Linux (not sure about others), calling shutdown+ -- on the write channel fails with with ENOTCONN if the+ -- write channel is already closed. It is common for this to+ -- happen (e.g. if the peer calls @close@ before the local+ -- process runs gracefulClose, the local operating system+ -- will have already closed the write channel). However,+ -- it does not pose a problem. We just proceed as we would+ -- have since either way we become certain that the write channel+ -- is closed.+ Left err -> if err == eNOTCONN+ then gracefulCloseB fd+ else do+ _ <- S.uninterruptibleClose fd+ -- TODO: What about ENOTCONN? Can this happen if the remote+ -- side has already closed the connection?+ throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ SCK.functionGracefulClose+ [SCK.cshutdown,describeErrorCode err]+ Right _ -> gracefulCloseB fd++gracefulCloseB :: Fd -> IO (Either CloseException ())+gracefulCloseB fd = do+ buf <- PM.newByteArray 1+ S.uninterruptibleReceiveMutableByteArray fd buf 0 1 mempty >>= \case+ Left err1 -> if err1 == eWOULDBLOCK || err1 == eAGAIN+ then do+ threadWaitRead fd+ -- TODO: We do not actually want to remove the bytes from the+ -- receive buffer. We should use MSG_PEEK instead. Then we will+ -- be certain to send a reset when a CloseException is reported.+ S.uninterruptibleReceiveMutableByteArray fd buf 0 1 mempty >>= \case+ Left err -> do+ _ <- S.uninterruptibleClose fd+ throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ SCK.functionGracefulClose+ [SCK.crecv,describeErrorCode err]+ Right sz -> if sz == 0+ then S.uninterruptibleClose fd >>= \case+ Left err -> throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ SCK.functionGracefulClose+ [SCK.cclose,describeErrorCode err]+ Right _ -> pure (Right ())+ else do+ debug ("Socket.Stream.IPv4.gracefulClose: remote not shutdown A")+ _ <- S.uninterruptibleClose fd+ pure (Left ClosePeerContinuedSending)+ else do+ _ <- S.uninterruptibleClose fd+ -- We treat all @recv@ errors except for the nonblocking+ -- notices as unrecoverable.+ throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ SCK.functionGracefulClose+ [SCK.crecv,describeErrorCode err1]+ Right sz -> if sz == 0+ then S.uninterruptibleClose fd >>= \case+ Left err -> throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ SCK.functionGracefulClose+ [SCK.cclose,describeErrorCode err]+ Right _ -> pure (Right ())+ else do+ debug ("Socket.Stream.IPv4.gracefulClose: remote not shutdown B")+ _ <- S.uninterruptibleClose fd+ pure (Left ClosePeerContinuedSending)++-- | Accept a connection on the listener and run the supplied callback+-- on it. This closes the connection when the callback finishes or if+-- an exception is thrown. Since this function blocks the thread until+-- the callback finishes, it is only suitable for stream socket clients+-- that handle one connection at a time. The variant 'forkAcceptedUnmasked'+-- is preferrable for servers that need to handle connections concurrently+-- (most use cases).+withAccepted ::+ Listener+ -> (Either CloseException () -> a -> IO b)+ -- ^ Callback to handle an ungraceful close. + -> (Connection -> Endpoint -> IO a)+ -- ^ Callback to consume connection. Must not return the connection.+ -> IO (Either (AcceptException 'Uninterruptible) b)+withAccepted lstn consumeException cb = do+ r <- mask $ \restore -> do+ accept lstn >>= \case+ Left e -> pure (Left e)+ Right (conn, endpoint) -> do+ a <- onException (restore (cb conn endpoint)) (disconnect_ conn)+ e <- disconnect conn+ pure (Right (e,a))+ -- Notice that consumeException gets run in an unmasked context.+ case r of+ Left e -> pure (Left e)+ Right (e,a) -> fmap Right (consumeException e a)++-- | Accept a connection on the listener and run the supplied callback in+-- a new thread. Prefer 'forkAcceptedUnmasked' unless the masking state+-- needs to be preserved for the callback. Such a situation seems unlikely+-- to the author.+forkAccepted ::+ Listener+ -> (Either CloseException () -> a -> IO ())+ -- ^ Callback to handle an ungraceful close. + -> (Connection -> Endpoint -> IO a)+ -- ^ Callback to consume connection. Must not return the connection.+ -> IO (Either (AcceptException 'Uninterruptible) ThreadId)+forkAccepted lst consumeException cb = internalAccepted+ ( \restore action -> do+ tid <- forkIO $ do+ (e,x) <- action restore+ restore (consumeException e x)+ pure (Right tid)+ ) lst cb++-- | Accept a connection on the listener and run the supplied callback in+-- a new thread. The masking state is set to @Unmasked@ when running the+-- callback. Typically, @a@ is instantiated to @()@.+forkAcceptedUnmasked ::+ Listener+ -> (Either CloseException () -> a -> IO ())+ -- ^ Callback to handle an ungraceful close. + -> (Connection -> Endpoint -> IO a)+ -- ^ Callback to consume connection. Must not return the connection.+ -> IO (Either (AcceptException 'Uninterruptible) ThreadId)+forkAcceptedUnmasked lstn consumeException cb =+ mask_ $ accept lstn >>= \case+ Left e -> pure (Left e)+ Right (conn, endpoint) -> fmap Right $ forkIOWithUnmask $ \unmask -> do+ a <- onException (unmask (cb conn endpoint)) (disconnect_ conn)+ e <- disconnect conn+ unmask (consumeException e a)++-- | Accept a connection on the listener and run the supplied callback in+-- a new thread. The masking state is set to @Unmasked@ when running the+-- callback. Typically, @a@ is instantiated to @()@.+--+-- ==== __Discussion__+--+-- Why is the @counter@ argument present? At first, it seems+-- like this is something that the API consumer should implement on+-- top of this library. The argument for the inclusion of the counter+-- is has two parts: (1) clients supporting graceful termination+-- always need these semantics and (2) these semantics cannot+-- be provided without building in @counter@ as a @TVar@.+--+-- 1. Clients supporting graceful termination always need these+-- semantics. To gracefully bring down a server that has been+-- accepting connections with a forking function, an application+-- must wait for all active connections to finish. Since all+-- connections run on separate threads, this can only be+-- accomplished by a concurrency primitive. The straightforward+-- solution is to wrap a counter with either @MVar@ or @TVar@.+-- To complete graceful termination, the application must+-- block until the counter reaches zero.+-- 2. These semantics cannot be provided without building in+-- @counter@ as a @TVar@. When @abandon@ becomes @True@,+-- graceful termination begins. From this point onward, if at+-- any point the counter reaches zero, the application consuming+-- this API will complete termination. Consequently, we need+-- the guarantee that the counter does not increment after+-- the @abandon@ transaction completes. If it did increment+-- in this forbidden way (e.g. if it was incremented some+-- unspecified amount of time after a connection was accepted),+-- there would be a race condition in which the application+-- may terminate without giving the newly accepted connection+-- a chance to finish. Fortunately, @STM@ gives us the+-- composable transaction we need to get this guarantee.+-- To wait for an inbound connection, we use:+--+-- > (isReady,deregister) <- threadWaitReadSTM fd+-- > shouldReceive <- atomically $ do+-- > readTVar abandon >>= \case+-- > True -> do+-- > isReady+-- > modifyTVar' counter (+1)+-- > pure True+-- > False -> pure False+--+-- This eliminates the window for the race condition. If a+-- connection is accepted, the counter is guaranteed to+-- be incremented _before_ @abandon@ becomes @True@.+-- However, this code would be more simple and would perform+-- better if GHC's event manager used TVar instead of STM.+interruptibleForkAcceptedUnmasked ::+ TVar Int+ -- ^ Connection counter. Incremented when connection+ -- is accepted. Decremented after connection is closed.+ -> TVar Bool+ -- ^ Interrupted. If this becomes 'True' give up and return+ -- @'Left' 'AcceptInterrupted'@.+ -> Listener+ -- ^ Connection listener+ -> (Either CloseException () -> a -> IO ())+ -- ^ Callback to handle an ungraceful close. + -> (Connection -> Endpoint -> IO a)+ -- ^ Callback to consume connection. Must not return the connection.+ -> IO (Either (AcceptException 'Interruptible) ThreadId)+interruptibleForkAcceptedUnmasked !counter !abandon !lstn consumeException cb =+ mask_ $ interruptibleAcceptCounting counter abandon lstn >>= \case+ Left e -> do+ case e of+ AcceptInterrupted -> pure ()+ _ -> atomically (modifyTVar' counter (subtract 1))+ pure (Left e)+ Right (conn, endpoint) -> fmap Right $ forkIOWithUnmask $ \unmask -> do+ a <- onException+ (unmask (cb conn endpoint))+ (disconnect_ conn *> atomically (modifyTVar' counter (subtract 1)))+ e <- disconnect conn+ r <- unmask (consumeException e a)+ atomically (modifyTVar' counter (subtract 1))+ pure r++-- | Open a socket and connect to a peer. Requirements:+--+-- * This function may only be called in contexts where exceptions+-- are masked.+-- * The caller /must/ be sure to call 'disconnect' or 'disconnect_'+-- on the resulting 'Connection' exactly once to close underlying+-- file descriptor.+-- * The 'Connection' cannot be used after being given as an argument+-- to 'disconnect' or 'disconnect_'.+--+-- Noncompliant use of this function leads to undefined behavior. Prefer+-- 'withConnection' unless you are writing an integration with a+-- resource-management library.+connect ::+ Endpoint+ -- ^ Remote endpoint+ -> IO (Either (ConnectException 'Uninterruptible) Connection)+connect !remote = do+ debug ("connect: opening connection " ++ show remote)+ e1 <- S.uninterruptibleSocket S.internet+ (L.applySocketFlags (L.closeOnExec <> L.nonblocking) S.stream)+ S.defaultProtocol+ debug ("connect: opened connection " ++ show remote)+ case e1 of+ Left err -> handleSocketConnectException SCK.functionWithConnection err+ Right fd -> do+ let sockAddr = id+ $ S.encodeSocketAddressInternet+ $ endpointToSocketAddressInternet+ $ remote+ merr <- S.uninterruptibleConnect fd sockAddr >>= \case+ Left err2 -> if err2 == eINPROGRESS+ then do+ threadWaitWrite fd+ pure Nothing+ else pure (Just err2)+ Right _ -> pure Nothing+ case merr of+ Just err -> do+ S.uninterruptibleErrorlessClose fd+ handleConnectException SCK.functionWithConnection err+ Nothing -> do+ e <- S.uninterruptibleGetSocketOption fd+ S.levelSocket S.optionError (intToCInt (PM.sizeOf (undefined :: CInt)))+ case e of+ Left err -> do+ S.uninterruptibleErrorlessClose fd+ throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ functionWithListener+ [SCK.cgetsockopt,describeEndpoint remote,describeErrorCode err]+ Right (sz,S.OptionValue val) -> if sz == intToCInt (PM.sizeOf (undefined :: CInt))+ then+ let err = PM.indexByteArray val 0 :: CInt in+ if err == 0+ then pure (Right (Connection fd))+ else do+ S.uninterruptibleErrorlessClose fd+ handleConnectException SCK.functionWithConnection (Errno err)+ else do+ S.uninterruptibleErrorlessClose fd+ throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ functionWithListener+ [SCK.cgetsockopt,describeEndpoint remote,connectErrorOptionValueSize]++-- | Close a connection gracefully, reporting a 'CloseException' when+-- the connection has to be terminated by sending a TCP reset. This+-- uses a combination of @shutdown@, @recv@, @close@ to detect when+-- resets need to be sent.+disconnect :: Connection -> IO (Either CloseException ())+disconnect (Connection fd) = gracefulCloseA fd++-- | Close a connection. This does not check to see whether or not+-- the connection was brought down gracefully. It just calls @close@+-- and is likely to cause a TCP reset to be sent. It never+-- throws exceptions of any kind (even if @close@ fails).+-- This should only be preferred+-- to 'disconnect' in exception-cleanup contexts where there is+-- already an exception that will be rethrown. See the implementation+-- of 'withConnection' for an example of appropriate use of both+-- 'disconnect' and 'disconnect_'.+disconnect_ :: Connection -> IO ()+disconnect_ (Connection fd) = S.uninterruptibleErrorlessClose fd++-- | Establish a connection to a server.+withConnection ::+ Endpoint+ -- ^ Remote endpoint+ -> (Either CloseException () -> a -> IO b)+ -- ^ Callback to handle an ungraceful close. + -> (Connection -> IO a)+ -- ^ Callback to consume connection. Must not return the connection.+ -> IO (Either (ConnectException 'Uninterruptible) b)+withConnection !remote g f = mask $ \restore -> do+ connect remote >>= \case+ Left err -> pure (Left err)+ Right conn -> do+ a <- onException (restore (f conn)) (disconnect_ conn)+ m <- disconnect conn+ b <- g m a+ pure (Right b)+ +sendByteArray ::+ Connection -- ^ Connection+ -> ByteArray -- ^ Payload+ -> IO (Either (SendException 'Uninterruptible) ())+sendByteArray conn arr =+ sendByteArraySlice conn arr 0 (PM.sizeofByteArray arr)++interruptibleSendByteArray ::+ TVar Bool+ -- ^ Interrupted. If this becomes 'True', give up and return+ -- @'Left' 'AcceptInterrupted'@.+ -> Connection -- ^ Connection+ -> ByteArray -- ^ Payload+ -> IO (Either (SendException 'Interruptible) ())+interruptibleSendByteArray abandon conn arr =+ interruptibleSendByteArraySlice abandon conn arr 0 (PM.sizeofByteArray arr)++sendByteArraySlice ::+ Connection -- ^ Connection+ -> ByteArray -- ^ Payload (will be sliced)+ -> Int -- ^ Offset into payload+ -> Int -- ^ Length of slice into buffer+ -> IO (Either (SendException 'Uninterruptible) ())+sendByteArraySlice !conn !payload !off0 !len0 = go off0 len0+ where+ go !off !len = if len > 0+ then internalSend conn payload off len >>= \case+ Left e -> pure (Left e)+ Right sz' -> do+ let sz = csizeToInt sz'+ go (off + sz) (len - sz)+ else if len == 0+ then pure (Right ())+ else throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ functionSendByteArray+ [SCK.negativeSliceLength]++interruptibleSendByteArraySlice ::+ TVar Bool+ -- ^ Interrupted. If this becomes 'True', give up and return+ -- @'Left' 'AcceptInterrupted'@.+ -> Connection -- ^ Connection+ -> ByteArray -- ^ Payload (will be sliced)+ -> Int -- ^ Offset into payload+ -> Int -- ^ Length of slice into buffer+ -> IO (Either (SendException 'Interruptible) ())+interruptibleSendByteArraySlice !abandon !conn !payload !off0 !len0 = go off0 len0+ where+ go !off !len = if len > 0+ then internalInterruptibleSend abandon conn payload off len >>= \case+ Left e -> pure (Left e)+ Right sz' -> do+ let sz = csizeToInt sz'+ go (off + sz) (len - sz)+ else if len == 0+ then pure (Right ())+ else throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ functionSendByteArray+ [SCK.negativeSliceLength]++sendMutableByteArray ::+ Connection -- ^ Connection+ -> MutableByteArray RealWorld -- ^ Buffer (will be sliced)+ -> IO (Either (SendException 'Uninterruptible) ())+sendMutableByteArray conn arr =+ sendMutableByteArraySlice conn arr 0 =<< PM.getSizeofMutableByteArray arr++sendMutableByteArraySlice ::+ Connection -- ^ Connection+ -> MutableByteArray RealWorld -- ^ Buffer (will be sliced)+ -> Int -- ^ Offset into payload+ -> Int -- ^ Length of slice into buffer+ -> IO (Either (SendException 'Uninterruptible) ())+sendMutableByteArraySlice !conn !payload !off0 !len0 = go off0 len0+ where+ go !off !len = if len > 0+ then internalSendMutable conn payload off len >>= \case+ Left e -> pure (Left e)+ Right sz' -> do+ let sz = csizeToInt sz'+ go (off + sz) (len - sz)+ else if len == 0+ then pure (Right ())+ else throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ functionSendMutableByteArray+ [SCK.negativeSliceLength]++interruptibleSendMutableByteArraySlice ::+ TVar Bool+ -- ^ Interrupted. If this becomes 'True' give up and return+ -- @'Left' 'AcceptInterrupted'@.+ -> Connection -- ^ Connection+ -> MutableByteArray RealWorld -- ^ Buffer (will be sliced)+ -> Int -- ^ Offset into payload+ -> Int -- ^ Length of slice into buffer+ -> IO (Either (SendException 'Interruptible) ())+interruptibleSendMutableByteArraySlice !abandon !conn !payload !off0 !len0 = go off0 len0+ where+ go !off !len = if len > 0+ then internalInterruptibleSendMutable abandon conn payload off len >>= \case+ Left e -> pure (Left e)+ Right sz' -> do+ let sz = csizeToInt sz'+ go (off + sz) (len - sz)+ else if len == 0+ then pure (Right ())+ else throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ functionSendMutableByteArray+ [SCK.negativeSliceLength]++-- Precondition: the length must be greater than zero.+internalInterruptibleSendMutable ::+ TVar Bool+ -- ^ Interrupted. If this becomes 'True' give up and return+ -- @'Left' 'AcceptInterrupted'@.+ -> Connection -- ^ Connection+ -> MutableByteArray RealWorld -- ^ Buffer (will be sliced)+ -> Int -- ^ Offset into payload+ -> Int -- ^ Length of slice into buffer+ -> IO (Either (SendException 'Interruptible) CSize)+internalInterruptibleSendMutable !abandon !conn !payload !off !len =+ veryInternalSendMutable+ (\fd -> interruptibleWaitWrite abandon fd >>= \case+ True -> pure (Right ())+ False -> pure (Left SendInterrupted)+ ) conn payload off len++-- Precondition: the length must be greater than zero.+internalSendMutable ::+ Connection -- ^ Connection+ -> MutableByteArray RealWorld -- ^ Buffer (will be sliced)+ -> Int -- ^ Offset into payload+ -> Int -- ^ Length of slice into buffer+ -> IO (Either (SendException 'Uninterruptible) CSize)+internalSendMutable !conn !payload !off !len =+ veryInternalSendMutable+ (\fd -> threadWaitWrite fd *> pure (Right ()))+ conn payload off len++-- Precondition: the length must be greater than zero.+veryInternalSendMutable :: + (Fd -> IO (Either (SendException i) ()))+ -> Connection -- ^ Connection+ -> MutableByteArray RealWorld -- ^ Buffer (will be sliced)+ -> Int -- ^ Offset into payload+ -> Int -- ^ Length of slice into buffer+ -> IO (Either (SendException i) CSize)+{-# INLINE veryInternalSendMutable #-}+veryInternalSendMutable wait (Connection !s) !payload !off !len = do+ e1 <- S.uninterruptibleSendMutableByteArray s payload+ (intToCInt off)+ (intToCSize len)+ (S.noSignal)+ case e1 of+ Left err1 -> if err1 == eWOULDBLOCK || err1 == eAGAIN+ then do+ wait s >>= \case+ Left err2 -> pure (Left err2)+ Right () -> do+ e3 <- S.uninterruptibleSendMutableByteArray s payload+ (intToCInt off)+ (intToCSize len)+ (S.noSignal)+ case e3 of+ Left err3 -> handleSendException functionSendMutableByteArray err3+ Right sz -> pure (Right sz)+ else handleSendException "sendMutableByteArray" err1+ Right sz -> pure (Right sz)++-- Precondition: the length must be greater than zero.+internalSend ::+ Connection -- ^ Connection+ -> ByteArray -- ^ Buffer (will be sliced)+ -> Int -- ^ Offset into payload+ -> Int -- ^ Length of slice into buffer+ -> IO (Either (SendException 'Uninterruptible) CSize)+internalSend !conn !payload !off !len = veryInternalSend+ (\fd -> threadWaitWrite fd *> pure (Right ()))+ conn payload off len++-- Precondition: the length must be greater than zero.+internalInterruptibleSend ::+ TVar Bool -- ^ Aliveness+ -> Connection -- ^ Connection+ -> ByteArray -- ^ Buffer (will be sliced)+ -> Int -- ^ Offset into payload+ -> Int -- ^ Length of slice into buffer+ -> IO (Either (SendException 'Interruptible) CSize)+internalInterruptibleSend !abandon !conn !payload !off !len = veryInternalSend+ (\fd -> interruptibleWaitWrite abandon fd >>= \case+ True -> pure (Right ())+ False -> pure (Left SendInterrupted)+ ) conn payload off len++-- Precondition: the length must be greater than zero.+veryInternalSend ::+ (Fd -> IO (Either (SendException i) ()))+ -> Connection -- ^ Connection+ -> ByteArray -- ^ Buffer (will be sliced)+ -> Int -- ^ Offset into payload+ -> Int -- ^ Length of slice into buffer+ -> IO (Either (SendException i) CSize)+{-# INLINE veryInternalSend #-}+veryInternalSend wait (Connection !s) !payload !off !len = do+ debug ("veryInternalSend: about to send chunk on stream socket, offset " ++ show off ++ " and length " ++ show len)+ e1 <- S.uninterruptibleSendByteArray s payload+ (intToCInt off)+ (intToCSize len)+ (S.noSignal)+ debug "veryInternalSend: just sent chunk on stream socket"+ case e1 of+ Left err1 -> if err1 == eWOULDBLOCK || err1 == eAGAIN+ then do+ debug "veryInternalSend: waiting to for write ready on stream socket"+ wait s >>= \case+ Left e -> pure (Left e)+ Right _ -> do+ e2 <- S.uninterruptibleSendByteArray s payload+ (intToCInt off)+ (intToCSize len)+ (S.noSignal)+ case e2 of+ Left err2 -> do+ debug "veryInternalSend: encountered error after sending chunk on stream socket"+ handleSendException functionSendByteArray err2+ Right sz -> pure (Right sz)+ else handleSendException functionSendByteArray err1+ Right sz -> pure (Right sz)++-- The maximum number of bytes to receive must be greater than zero.+-- The operating system guarantees us that the returned actual number+-- of bytes is less than or equal to the requested number of bytes.+-- This function does not validate that the result size is greater+-- than zero. Functions calling this must perform that check. This+-- also does not trim the buffer. The caller must do that if it is+-- necessary. This function does use the event manager to wait+-- for the socket to be ready for reads.+internalReceiveMaximally ::+ Connection -- ^ Connection+ -> Int -- ^ Maximum number of bytes to receive+ -> MutableByteArray RealWorld -- ^ Receive buffer+ -> Int -- ^ Offset into buffer+ -> IO (Either (ReceiveException i) Int)+internalReceiveMaximally (Connection !fd) !maxSz !buf !off = do+ debug "receive: stream socket about to wait"+ threadWaitRead fd+ debug ("receive: stream socket is now readable, receiving up to " ++ show maxSz ++ " bytes at offset " ++ show off)+ e <- S.uninterruptibleReceiveMutableByteArray fd buf (intToCInt off) (intToCSize maxSz) mempty+ debug "receive: finished reading from stream socket"+ case e of+ Left err -> handleReceiveException "internalReceiveMaximally" err+ Right recvSz -> pure (Right (csizeToInt recvSz))++internalInterruptibleReceiveMaximally ::+ TVar Bool -- ^ If this completes, give up on receiving+ -> Connection -- ^ Connection+ -> Int -- ^ Maximum number of bytes to receive+ -> MutableByteArray RealWorld -- ^ Receive buffer+ -> Int -- ^ Offset into buffer+ -> IO (Either (ReceiveException 'Interruptible) Int)+{-# INLINE internalInterruptibleReceiveMaximally #-}+internalInterruptibleReceiveMaximally abandon (Connection !fd) !maxSz !buf !off = do+ shouldReceive <- interruptibleWaitRead abandon fd+ if shouldReceive+ then do+ e <- S.uninterruptibleReceiveMutableByteArray fd buf (intToCInt off) (intToCSize maxSz) mempty+ case e of+ Left err -> handleReceiveException "internalReceiveMaximally" err+ Right recvSz -> pure (Right (csizeToInt recvSz))+ else pure (Left ReceiveInterrupted)++-- | Receive exactly the given number of bytes. If the remote application+-- shuts down its end of the connection before sending the required+-- number of bytes, this returns @'Left' 'ReceiveShutdown'@.+receiveByteArray ::+ Connection -- ^ Connection+ -> Int -- ^ Number of bytes to receive+ -> IO (Either (ReceiveException 'Uninterruptible) ByteArray)+receiveByteArray !conn !total =+ internalReceiveByteArray internalReceiveMaximally conn total++-- | Variant of 'receiveByteArray' that support STM-style interrupts.+interruptibleReceiveByteArray ::+ TVar Bool+ -- ^ Interrupted. If this becomes 'True' give up and return+ -- @'Left' 'ReceiveInterrupted'@.+ -> Connection -- ^ Connection+ -> Int -- ^ Number of bytes to receive+ -> IO (Either (ReceiveException 'Interruptible) ByteArray)+interruptibleReceiveByteArray !abandon !conn !total =+ internalReceiveByteArray (internalInterruptibleReceiveMaximally abandon) conn total++-- This is used by both receiveByteArray and interruptibleReceiveByteArray.+internalReceiveByteArray ::+ (Connection -> Int -> MutableByteArray RealWorld -> Int -> IO (Either (ReceiveException i) Int))+ -> Connection+ -> Int+ -> IO (Either (ReceiveException i) ByteArray)+internalReceiveByteArray recvMax !conn0 !total = do+ marr <- PM.newByteArray total+ go conn0 marr 0 total+ where+ go !conn !marr !off !remaining = case compare remaining 0 of+ GT -> do+ recvMax conn remaining marr off >>= \case+ Left err -> pure (Left err)+ Right sz -> if sz /= 0+ then go conn marr (off + sz) (remaining - sz)+ else pure (Left ReceiveShutdown)+ EQ -> do+ arr <- PM.unsafeFreezeByteArray marr+ pure (Right arr)+ LT -> throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ functionReceiveByteArray+ [SCK.negativeSliceLength]++-- | Receive a number of bytes exactly equal to the size of the mutable+-- byte array. If the remote application shuts down its end of the+-- connection before sending the required number of bytes, this returns+-- @'Left' ('SocketException' 'Receive' 'RemoteShutdown')@.+receiveMutableByteArray ::+ Connection+ -> MutableByteArray RealWorld+ -> IO (Either (ReceiveException 'Uninterruptible) ())+receiveMutableByteArray !conn0 !marr0 = do+ total <- PM.getSizeofMutableByteArray marr0+ go conn0 marr0 0 total+ where+ go !conn !marr !off !remaining = if remaining > 0+ then do+ internalReceiveMaximally conn remaining marr off >>= \case+ Left err -> pure (Left err)+ Right sz -> if sz /= 0+ then go conn marr (off + sz) (remaining - sz)+ else pure (Left ReceiveShutdown)+ else pure (Right ())++-- | Receive up to the given number of bytes, using the given array and+-- starting at the given offset.+receiveBoundedMutableByteArraySlice ::+ Connection -- ^ Connection+ -> Int -- ^ Maximum number of bytes to receive+ -> MutableByteArray RealWorld -- ^ Buffer in which the data are going to be stored+ -> Int -- ^ Offset in the buffer+ -> IO (Either (ReceiveException 'Uninterruptible) Int) -- ^ Either a socket exception or the number of bytes read+receiveBoundedMutableByteArraySlice !conn !total !marr !off+ | total > 0 = do+ internalReceiveMaximally conn total marr off >>= \case+ Left err -> pure (Left err)+ Right sz -> if sz /= 0+ then pure (Right sz)+ else pure (Left ReceiveShutdown)+ | total == 0 = pure (Right 0)+ | otherwise = throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ functionReceiveMutableByteArraySlice+ [SCK.negativeSliceLength]++-- | Receive up to the given number of bytes, using the given array and+-- starting at the given offset. This can be interrupted by the+-- completion of an 'STM' transaction.+interruptibleReceiveBoundedMutableByteArraySlice ::+ TVar Bool+ -- ^ Interrupted. If this becomes 'True' give up and return+ -- @'Left' 'ReceiveInterrupted'@.+ -> Connection -- ^ Connection+ -> Int -- ^ Maximum number of bytes to receive+ -> MutableByteArray RealWorld -- ^ Buffer in which the data are going to be stored+ -> Int -- ^ Offset in the buffer+ -> IO (Either (ReceiveException 'Interruptible) Int) -- ^ Either a socket exception or the number of bytes read+interruptibleReceiveBoundedMutableByteArraySlice !abandon !conn !total !marr !off+ | total > 0 = do+ internalInterruptibleReceiveMaximally abandon conn total marr off >>= \case+ Left err -> pure (Left err)+ Right sz -> if sz /= 0+ then pure (Right sz)+ else pure (Left ReceiveShutdown)+ | total == 0 = pure (Right 0)+ | otherwise = throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ -- TODO: fix this function name in the error reporting+ functionReceiveMutableByteArraySlice+ [SCK.negativeSliceLength]++-- | Receive up to the given number of bytes. If the remote application+-- shuts down its end of the connection instead of sending any bytes,+-- this returns+-- @'Left' ('SocketException' 'Receive' 'RemoteShutdown')@.+receiveBoundedByteArray ::+ Connection -- ^ Connection+ -> Int -- ^ Maximum number of bytes to receive+ -> IO (Either (ReceiveException 'Uninterruptible) ByteArray)+receiveBoundedByteArray !conn !total+ | total > 0 = do+ m <- PM.newByteArray total+ receiveBoundedMutableByteArraySlice conn total m 0 >>= \case+ Left err -> pure (Left err)+ Right sz -> do+ shrinkMutableByteArray m sz+ Right <$> PM.unsafeFreezeByteArray m+ | total == 0 = pure (Right mempty)+ | otherwise = throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ functionReceiveBoundedByteArray+ [SCK.negativeSliceLength]++endpointToSocketAddressInternet :: Endpoint -> S.SocketAddressInternet+endpointToSocketAddressInternet (Endpoint {address, port}) = S.SocketAddressInternet+ { port = S.hostToNetworkShort port+ , address = S.hostToNetworkLong (getIPv4 address)+ }++socketAddressInternetToEndpoint :: S.SocketAddressInternet -> Endpoint+socketAddressInternetToEndpoint (S.SocketAddressInternet {address,port}) = Endpoint+ { address = IPv4 (S.networkToHostLong address)+ , port = S.networkToHostShort port+ }++intToCInt :: Int -> CInt+intToCInt = fromIntegral++intToCSize :: Int -> CSize+intToCSize = fromIntegral++csizeToInt :: CSize -> Int+csizeToInt = fromIntegral++shrinkMutableByteArray :: MutableByteArray RealWorld -> Int -> IO ()+shrinkMutableByteArray (MutableByteArray arr) (I# sz) =+ PM.primitive_ (shrinkMutableByteArray# arr sz)++moduleSocketStreamIPv4 :: String+moduleSocketStreamIPv4 = "Socket.Stream.IPv4"++functionSendMutableByteArray :: String+functionSendMutableByteArray = "sendMutableByteArray"++functionSendByteArray :: String+functionSendByteArray = "sendByteArray"++functionWithListener :: String+functionWithListener = "withListener"++functionReceiveBoundedByteArray :: String+functionReceiveBoundedByteArray = "receiveBoundedByteArray"++functionReceiveByteArray :: String+functionReceiveByteArray = "receiveByteArray"++functionReceiveMutableByteArraySlice :: String+functionReceiveMutableByteArraySlice = "receiveMutableByteArraySlice"++describeErrorCode :: Errno -> String+describeErrorCode (Errno e) = "error code " ++ show e++handleReceiveException :: String -> Errno -> IO (Either (ReceiveException i) a)+handleReceiveException func e+ | e == eCONNRESET = pure (Left ReceiveReset)+ | otherwise = throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ func+ [describeErrorCode e]++handleSendException :: String -> Errno -> IO (Either (SendException i) a)+{-# INLINE handleSendException #-}+handleSendException func e+ | e == ePIPE = pure (Left SendShutdown)+ | e == eCONNRESET = pure (Left SendReset)+ | otherwise = throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ func+ [describeErrorCode e]++-- These are the exceptions that can happen as a result of a+-- nonblocking @connect@ or as a result of subsequently calling+-- @getsockopt@ to get the @SO_ERROR@ after the socket is ready+-- for writes. For increased likelihood of correctness, we do+-- not distinguish between which of these error codes can show+-- up after which of those two calls. This means we are likely+-- testing for some exceptions that cannot occur as a result of+-- a particular call, but doing otherwise would be fraught with+-- uncertainty.+handleConnectException :: String -> Errno -> IO (Either (ConnectException i) a)+handleConnectException func e+ | e == eACCES = pure (Left ConnectFirewalled)+ | e == ePERM = pure (Left ConnectFirewalled)+ | e == eNETUNREACH = pure (Left ConnectNetworkUnreachable)+ | e == eCONNREFUSED = pure (Left ConnectRefused)+ | e == eADDRNOTAVAIL = pure (Left ConnectEphemeralPortsExhausted)+ | e == eTIMEDOUT = pure (Left ConnectTimeout)+ | otherwise = throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ func+ [describeErrorCode e]++-- These are the exceptions that can happen as a result+-- of calling @socket@ with the intent of using the socket+-- to open a connection (not listen for inbound connections).+handleSocketConnectException :: String -> Errno -> IO (Either (ConnectException i) a)+handleSocketConnectException func e+ | e == eMFILE = pure (Left ConnectFileDescriptorLimit)+ | e == eNFILE = pure (Left ConnectFileDescriptorLimit)+ | otherwise = throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ func+ [describeErrorCode e]++-- These are the exceptions that can happen as a result+-- of calling @socket@ with the intent of using the socket+-- to listen for inbound connections.+handleSocketListenException :: String -> Errno -> IO (Either SocketException a)+handleSocketListenException func e+ | e == eMFILE = pure (Left SocketFileDescriptorLimit)+ | e == eNFILE = pure (Left SocketFileDescriptorLimit)+ | otherwise = throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ func+ [describeErrorCode e]++-- These are the exceptions that can happen as a result+-- of calling @bind@ with the intent of using the socket+-- to listen for inbound connections. This is also used+-- to clean up the error codes of @listen@. The two can+-- report some of the same error codes, and those happen+-- to be the error codes we are interested in.+--+-- NB: EACCES only happens on @bind@, not on @listen@.+handleBindListenException :: Word16 -> String -> Errno -> IO (Either SocketException a)+handleBindListenException thePort func e+ | e == eACCES = pure (Left SocketPermissionDenied)+ | e == eADDRINUSE = if thePort == 0+ then pure (Left SocketAddressInUse)+ else pure (Left SocketEphemeralPortsExhausted)+ | otherwise = throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ func+ [describeErrorCode e]++-- These are the exceptions that can happen as a result+-- of calling @socket@ with the intent of using the socket+-- to open a connection (not listen for inbound connections).+handleAcceptException :: String -> Errno -> IO (Either (AcceptException i) a)+handleAcceptException func e+ | e == eCONNABORTED = pure (Left AcceptConnectionAborted)+ | e == eMFILE = pure (Left AcceptFileDescriptorLimit)+ | e == eNFILE = pure (Left AcceptFileDescriptorLimit)+ | e == ePERM = pure (Left AcceptFirewalled)+ | otherwise = throwIO $ SocketUnrecoverableException+ moduleSocketStreamIPv4+ func+ [describeErrorCode e]++connectErrorOptionValueSize :: String+connectErrorOptionValueSize = "incorrectly sized value of SO_ERROR option"++interruptibleWaitRead :: TVar Bool -> Fd -> IO Bool+interruptibleWaitRead !abandon !fd = do+ (isReady,deregister) <- threadWaitReadSTM fd+ shouldReceive <- atomically+ ((bool retry (pure False) =<< readTVar abandon) <|> (isReady $> True))+ deregister+ pure shouldReceive++interruptibleWaitWrite :: TVar Bool -> Fd -> IO Bool+interruptibleWaitWrite !abandon !fd = do+ (isReady,deregister) <- threadWaitWriteSTM fd+ shouldSend <- atomically+ ((bool retry (pure False) =<< readTVar abandon) <|> (isReady $> True))+ deregister+ pure shouldSend++interruptibleWaitReadCounting :: TVar Int -> TVar Bool -> Fd -> IO Bool+interruptibleWaitReadCounting !counter !abandon !fd = do+ (isReady,deregister) <- threadWaitReadSTM fd+ shouldReceive <- atomically $ do+ readTVar abandon >>= \case+ False -> do+ isReady+ modifyTVar' counter (+1)+ pure True+ True -> pure False+ deregister+ pure shouldReceive++{- $unbracketed+ +Provided here are the unbracketed functions for the creation and destruction+of listeners, outbound connections, and inbound connections. These functions+come with pretty serious requirements:++* They may only be called in contexts where exceptions are masked.+* The caller /must/ be sure to call the destruction function every+ 'Listener' or 'Connection' exactly once to close underlying file+ descriptor.+* The 'Listener' or 'Connection' cannot be used after being given+ as an argument to the destruction function.+-}+
test/Main.hs view
@@ -1,38 +1,68 @@ {-# language BangPatterns #-}+{-# language DerivingStrategies #-}+{-# language DeriveAnyClass #-}+{-# language LambdaCase #-} {-# language ScopedTypeVariables #-}+{-# language TypeFamilies #-} import Control.Concurrent.Async (concurrently)+import Control.Monad (replicateM_) import Control.Exception (Exception) import Control.Exception (throwIO) import Control.Monad.ST (runST)+import Data.Bool (bool) import Data.Primitive (ByteArray) import Data.Word (Word16,Word8) import GHC.Exts (RealWorld)+import System.Exit (exitFailure)+import System.IO (stderr,hPutStrLn) import Test.Tasty import Test.Tasty.HUnit -import qualified Socket.Datagram.IPv4.Undestined as DIU-import qualified Socket.Stream.IPv4 as SI-import qualified GHC.Exts as E import qualified Data.Primitive as PM import qualified Data.Primitive.MVar as PM+import qualified GHC.Exts as E import qualified Net.IPv4 as IPv4+import qualified Socket.Datagram.IPv4.Spoof as DIS+import qualified Socket.Datagram.IPv4.Undestined as DIU+import qualified Socket.Stream.IPv4 as SI main :: IO ()-main = defaultMain tests+main = do+ canSpoof <- DIS.withSocket (const (pure ())) >>= \case+ Right () -> pure True+ Left e -> case e of+ DIS.SocketPermissionDenied -> pure False+ DIS.SocketFileDescriptorLimit -> do+ hPutStrLn stderr "All ephemeral ports are in use. Terminating."+ exitFailure+ defaultMain (tests canSpoof) -tests :: TestTree-tests = testGroup "socket"+tests :: Bool -> TestTree+tests canSpoof = testGroup "socket" [ testGroup "datagram" [ testGroup "ipv4" [ testGroup "undestined" [ testCase "A" testDatagramUndestinedA+ , testCase "B" testDatagramUndestinedB+ , testCase "C" testDatagramUndestinedC ]+ , testGroup "spoof" $ if canSpoof+ then+ [ testCase "A" testDatagramSpoofA+ , testCase "B" testDatagramSpoofB+ ]+ else [] ] ] , testGroup "stream" [ testGroup "ipv4" [ testCase "A" testStreamA+ , testGroup "B"+ [ testCase "1MB" (testStreamB 1)+ , testCase "4MB" (testStreamB 4)+ , testCase "32MB" (testStreamB 32)+ ] ] ] ]@@ -40,11 +70,24 @@ unhandled :: Exception e => IO (Either e a) -> IO a unhandled action = action >>= either throwIO pure +unhandledClose :: Either SI.CloseException () -> a -> IO a+unhandledClose m a = case m of+ Right () -> pure a+ Left e -> throwIO e++data MagicByteMismatch = MagicByteMismatch+ deriving stock (Show,Eq)+ deriving anyclass (Exception)++data NegativeByteCount = NegativeByteCount+ deriving stock (Show,Eq)+ deriving anyclass (Exception)+ testDatagramUndestinedA :: Assertion testDatagramUndestinedA = do (m :: PM.MVar RealWorld Word16) <- PM.newEmptyMVar (port,received) <- concurrently (sender m) (receiver m)- received @=? (DIU.Endpoint IPv4.loopback port, message)+ received @=? DIU.Message (DIU.Endpoint IPv4.loopback port) message where message = E.fromList [0,1,2,3] :: ByteArray sz = PM.sizeofByteArray message@@ -53,11 +96,81 @@ dstPort <- PM.takeMVar m unhandled $ DIU.send sock (DIU.Endpoint IPv4.loopback dstPort) message 0 sz pure srcPort- receiver :: PM.MVar RealWorld Word16 -> IO (DIU.Endpoint,ByteArray)+ receiver :: PM.MVar RealWorld Word16 -> IO DIU.Message receiver m = unhandled $ DIU.withSocket (DIU.Endpoint IPv4.loopback 0) $ \sock port -> do PM.putMVar m port- unhandled $ DIU.receive sock sz+ unhandled $ DIU.receiveByteArray sock sz +testDatagramUndestinedB :: Assertion+testDatagramUndestinedB = do+ (m :: PM.MVar RealWorld Word16) <- PM.newEmptyMVar+ (n :: PM.MVar RealWorld ()) <- PM.newEmptyMVar+ (port,received) <- concurrently (sender m n) (receiver m n)+ received @=?+ ( DIU.Message (DIU.Endpoint IPv4.loopback port) message1+ , DIU.Message (DIU.Endpoint IPv4.loopback port) message2+ )+ where+ message1 = E.fromList [0,1,2,3] :: ByteArray+ message2 = E.fromList [4,5,6,8,9,10] :: ByteArray+ sz1 = PM.sizeofByteArray message1+ sz2 = PM.sizeofByteArray message2+ sender :: PM.MVar RealWorld Word16 -> PM.MVar RealWorld () -> IO Word16+ sender m n = unhandled $ DIU.withSocket (DIU.Endpoint IPv4.loopback 0) $ \sock srcPort -> do+ dstPort <- PM.takeMVar m+ unhandled $ DIU.send sock (DIU.Endpoint IPv4.loopback dstPort) message1 0 sz1+ unhandled $ DIU.send sock (DIU.Endpoint IPv4.loopback dstPort) message2 0 sz2+ PM.putMVar n ()+ pure srcPort+ receiver :: PM.MVar RealWorld Word16 -> PM.MVar RealWorld () -> IO (DIU.Message,DIU.Message)+ receiver m n = unhandled $ DIU.withSocket (DIU.Endpoint IPv4.loopback 0) $ \sock port -> do+ PM.putMVar m port+ PM.takeMVar n+ msgs <- unhandled $ DIU.receiveMany sock 3 (max sz1 sz2)+ if PM.sizeofArray msgs == 2+ then pure (PM.indexArray msgs 0, PM.indexArray msgs 1)+ else fail "received a number of messages other than 2"+ +testDatagramUndestinedC :: Assertion+testDatagramUndestinedC = do+ (m :: PM.MVar RealWorld Word16) <- PM.newEmptyMVar+ (n :: PM.MVar RealWorld ()) <- PM.newEmptyMVar+ (port,received) <- concurrently (sender m n) (receiver m n)+ received @=?+ ( DIU.Message (DIU.Endpoint IPv4.loopback port) message1+ , DIU.Message (DIU.Endpoint IPv4.loopback port) message2+ , DIU.Message (DIU.Endpoint IPv4.loopback port) message3+ )+ where+ message1 = E.fromList (enumFromTo 0 9):: ByteArray+ message2 = E.fromList (enumFromTo 10 10) :: ByteArray+ message3 = E.fromList (enumFromTo 11 25) :: ByteArray+ sz1 = PM.sizeofByteArray message1+ sz2 = PM.sizeofByteArray message2+ sz3 = PM.sizeofByteArray message3+ sender :: PM.MVar RealWorld Word16 -> PM.MVar RealWorld () -> IO Word16+ sender m n = unhandled $ DIU.withSocket (DIU.Endpoint IPv4.loopback 0) $ \sock srcPort -> do+ dstPort <- PM.takeMVar m+ unhandled $ DIU.send sock (DIU.Endpoint IPv4.loopback dstPort) message1 0 sz1+ unhandled $ DIU.send sock (DIU.Endpoint IPv4.loopback dstPort) message2 0 sz2+ unhandled $ DIU.send sock (DIU.Endpoint IPv4.loopback dstPort) message3 0 sz3+ PM.putMVar n ()+ pure srcPort+ receiver :: PM.MVar RealWorld Word16 -> PM.MVar RealWorld () -> IO (DIU.Message,DIU.Message,DIU.Message)+ receiver m n = unhandled $ DIU.withSocket (DIU.Endpoint IPv4.loopback 0) $ \sock port -> do+ PM.putMVar m port+ PM.takeMVar n+ msgsX <- unhandled $ DIU.receiveMany sock 2 (max sz1 sz2)+ (msg1,msg2) <- if PM.sizeofArray msgsX == 2+ then pure (PM.indexArray msgsX 0, PM.indexArray msgsX 1)+ else fail "received a number of messages other than 2"+ msgsY <- unhandled $ DIU.receiveMany sock 2 sz3+ msg3 <- if PM.sizeofArray msgsY == 1+ then pure (PM.indexArray msgsY 0)+ else fail "received a number of messages other than 2"+ pure (msg1,msg2,msg3)+ + -- This test involves a made up protocol that goes like this: -- The sender always starts by sending the length of the rest -- of the payload as a native-endian encoded machine-sized int.@@ -80,16 +193,126 @@ sender :: PM.MVar RealWorld Word16 -> IO () sender m = do dstPort <- PM.takeMVar m- unhandled $ SI.withConnection (DIU.Endpoint IPv4.loopback dstPort) $ \conn -> do+ unhandled $ SI.withConnection (DIU.Endpoint IPv4.loopback dstPort) unhandledClose $ \conn -> do unhandled $ SI.sendByteArray conn szb unhandled $ SI.sendByteArray conn message receiver :: PM.MVar RealWorld Word16 -> IO ByteArray receiver m = unhandled $ SI.withListener (SI.Endpoint IPv4.loopback 0) $ \listener port -> do PM.putMVar m port- unhandled $ SI.withAccepted listener $ \conn _ -> do+ unhandled $ SI.withAccepted listener unhandledClose $ \conn _ -> do serializedSize <- unhandled $ SI.receiveByteArray conn (PM.sizeOf (undefined :: Int)) let theSize = PM.indexByteArray serializedSize 0 :: Int result <- unhandled $ SI.receiveByteArray conn theSize pure result +-- The sender sends a large amount of traffic that may exceed+-- the size of the operating system's TCP send buffer. The +-- amount is configurable because the test suite wants to+-- check this for several values.+testStreamB :: Int -> Assertion+testStreamB megabytes = do+ (m :: PM.MVar RealWorld Word16) <- PM.newEmptyMVar+ ((),()) <- concurrently (sender m) (receiver m)+ pure ()+ where+ message = E.fromList (replicate (32 * 1024) magicByte) :: ByteArray+ chunkSize = PM.sizeofByteArray message+ sender :: PM.MVar RealWorld Word16 -> IO ()+ sender m = do+ dstPort <- PM.takeMVar m+ unhandled $ SI.withConnection (DIU.Endpoint IPv4.loopback dstPort) unhandledClose $ \conn -> do+ replicateM_ (32 * megabytes) $ unhandled $ SI.sendByteArray conn message+ receiver :: PM.MVar RealWorld Word16 -> IO ()+ receiver m = unhandled $ SI.withListener (SI.Endpoint IPv4.loopback 0) $ \listener port -> do+ PM.putMVar m port+ unhandled $ SI.withAccepted listener unhandledClose $ \conn _ -> do+ buffer <- PM.newByteArray chunkSize+ let receiveLoop !remaining+ | remaining > 0 = do+ let recvSize = min remaining chunkSize+ PM.setByteArray buffer 0 chunkSize (0 :: Word8)+ bytesReceived <- unhandled (SI.receiveBoundedMutableByteArraySlice conn recvSize buffer 0)+ verifyClientSendBytes buffer bytesReceived >>= \case+ True -> receiveLoop (remaining - bytesReceived)+ False -> throwIO MagicByteMismatch+ | remaining == 0 = pure ()+ | otherwise = throwIO NegativeByteCount+ receiveLoop (32 * megabytes * chunkSize)+ pure () +magicByte :: Word8+magicByte = 0xFA++verifyClientSendBytes :: PM.MutableByteArray RealWorld -> Int -> IO Bool+verifyClientSendBytes arr len = go (len - 1)+ where+ go !ix = if ix >= 0+ then do+ w <- PM.readByteArray arr ix+ if w == magicByte then go (ix - 1) else pure False+ else pure True++-- Here, the sender spoofs its ip address and port.+testDatagramSpoofA :: Assertion+testDatagramSpoofA = do+ (m :: PM.MVar RealWorld Word16) <- PM.newEmptyMVar+ ((),received) <- concurrently (sender m) (receiver m)+ received @=? DIU.Message+ (DIU.Endpoint (IPv4.fromOctets 8 7 6 5) 60000)+ payload+ where+ sz = 16+ payload = E.fromList (enumFromTo (0 :: Word8) (fromIntegral sz - 1))+ sender :: PM.MVar RealWorld Word16 -> IO ()+ sender m = unhandled $ DIS.withSocket $ \sock -> do+ dstPort <- PM.takeMVar m+ marr <- PM.newByteArray sz+ PM.copyByteArray marr 0 payload 0 sz+ unhandled $ DIS.sendMutableByteArray sock+ (DIU.Endpoint (IPv4.fromOctets 8 7 6 5) 60000)+ (DIU.Endpoint IPv4.loopback dstPort)+ marr 0 sz+ receiver :: PM.MVar RealWorld Word16 -> IO DIU.Message+ receiver m = unhandled $ DIU.withSocket (DIU.Endpoint IPv4.loopback 0) $ \sock port -> do+ PM.putMVar m port+ unhandled $ DIU.receiveByteArray sock 500+ +-- Here, the sender spoofs its ip address and port twice, picking a+-- different port each time.+testDatagramSpoofB :: Assertion+testDatagramSpoofB = do+ (m :: PM.MVar RealWorld Word16) <- PM.newEmptyMVar+ ((),received) <- concurrently (sender m) (receiver m)+ received @=?+ ( DIU.Message+ (DIU.Endpoint (IPv4.fromOctets 8 7 6 5) 60000)+ payloadA+ , DIU.Message+ (DIU.Endpoint (IPv4.fromOctets 9 8 7 6) 59999)+ payloadB+ )+ where+ sz = 16+ payloadA = E.fromList (enumFromTo (1 :: Word8) (fromIntegral sz))+ payloadB = E.fromList (enumFromTo (2 :: Word8) (fromIntegral sz + 1))+ sender :: PM.MVar RealWorld Word16 -> IO ()+ sender m = unhandled $ DIS.withSocket $ \sock -> do+ dstPort <- PM.takeMVar m+ marrA <- PM.newByteArray sz+ marrB <- PM.newByteArray sz+ PM.copyByteArray marrA 0 payloadA 0 sz+ PM.copyByteArray marrB 0 payloadB 0 sz+ unhandled $ DIS.sendMutableByteArray sock+ (DIU.Endpoint (IPv4.fromOctets 8 7 6 5) 60000)+ (DIU.Endpoint IPv4.loopback dstPort)+ marrA 0 sz+ unhandled $ DIS.sendMutableByteArray sock+ (DIU.Endpoint (IPv4.fromOctets 9 8 7 6) 59999)+ (DIU.Endpoint IPv4.loopback dstPort)+ marrB 0 sz+ receiver :: PM.MVar RealWorld Word16 -> IO (DIU.Message,DIU.Message)+ receiver m = unhandled $ DIU.withSocket (DIU.Endpoint IPv4.loopback 0) $ \sock port -> do+ PM.putMVar m port+ msg1 <- unhandled $ DIU.receiveByteArray sock 500+ msg2 <- unhandled $ DIU.receiveByteArray sock 500+ return (msg1,msg2)