capnp 0.14.0.0 → 0.15.0.0
raw patch · 17 files changed
+761/−509 lines, 17 files
Files
- CHANGELOG.md +10/−0
- bench/Main.hs +7/−1
- capnp.cabal +3/−2
- lib/Capnp/Address.hs +7/−0
- lib/Capnp/Canonicalize.hs +43/−93
- lib/Capnp/Message.hs +40/−35
- lib/Capnp/Mutability.hs +1/−1
- lib/Capnp/New/Rpc/Server.hs +2/−2
- lib/Capnp/Repr.hs +71/−8
- lib/Capnp/Repr/Methods.hs +18/−11
- lib/Capnp/Rpc/Transport.hs +34/−6
- lib/Capnp/Rpc/Untyped.hs +268/−201
- lib/Capnp/Untyped.hs +184/−120
- lib/Internal/BuildPure.hs +6/−27
- lib/Internal/STE.hs +65/−0
- tests/Module/Capnp/Canonicalize.hs +1/−1
- tests/WalkSchemaCodeGenRequest.hs +1/−1
CHANGELOG.md view
@@ -1,3 +1,13 @@+# 0.15.0.0++- There is now a per-connection limit on the total size of incoming+ `Call` messages that are being serviced, which can be used to+ limit memory usage and provide backpressure. `ConnConfig` has+ a new `maxCallWords` field to configure this.+- Some bugs in the RPC layer have been fixed.+- `tracingTransport` now provides an option to omit call & return+ bodies from the logged messages.+ # 0.14.0.0 - Significant performance improvements.
bench/Main.hs view
@@ -1,5 +1,6 @@ {-# LANGUAGE OverloadedStrings #-} {-# LANGUAGE QuasiQuotes #-}+{-# LANGUAGE TypeApplications #-} module Main (main) where @@ -42,9 +43,14 @@ msg <- C.bsToMsg cgrBytes let whnfLTIO = whnfIO . C.evalLimitT maxBound defaultMain- [ bench "canonicalize" $ whnfLTIO $ do+ [ bench "canonicalize/IO" $ whnfLTIO $ do root <- U.rootPtr msg C.canonicalize root+ , bench "canonicalize/PureBuilder" $ whnfLTIO $ do+ C.createPure maxBound $ do+ root <- U.rootPtr msg+ (msg, _seg) <- C.canonicalize root+ pure msg , env (C.evalLimitT maxBound $ do mutMsg <- thaw msg
capnp.cabal view
@@ -1,6 +1,6 @@ cabal-version: 2.2 name: capnp-version: 0.14.0.0+version: 0.15.0.0 category: Data, Serialization, Network, Rpc copyright: 2016-2021 haskell-capnp contributors (see CONTRIBUTORS file). author: Ian Denhardt@@ -53,7 +53,7 @@ common shared-opts build-depends:- base >= 4.11 && < 5+ base >= 4.14 && < 5 , bytes >= 0.15.4 && <0.18 , bytestring >= 0.10 && <0.12 , containers >= 0.5.9 && <0.7@@ -152,6 +152,7 @@ , Internal.Rc , Internal.TCloseQ , Internal.BuildPure+ , Internal.STE -- other-extensions: build-depends: hashable >= 1.2.7 && <1.4
lib/Capnp/Address.hs view
@@ -15,10 +15,12 @@ , OffsetError(..) , computeOffset , pointerFrom+ , resolveOffset ) where import Data.Bits+import Data.Int import Data.Word import Capnp.Bits (WordCount)@@ -85,3 +87,8 @@ pointerFrom ptrAddr targetAddr (P.ListPtr _ eltSpec) = flip fmap (computeOffset ptrAddr targetAddr) $ \off -> P.ListPtr (fromIntegral off) eltSpec++-- | Add an offset to a WordAddr.+resolveOffset :: WordAddr -> Int32 -> WordAddr+resolveOffset addr@WordAt{..} off =+ addr { wordIndex = wordIndex + fromIntegral off + 1 }
lib/Capnp/Canonicalize.hs view
@@ -8,6 +8,7 @@ {-# LANGUAGE TypeFamilies #-} module Capnp.Canonicalize ( canonicalize+ , canonicalizeMut ) where -- Note [Allocation strategy]@@ -34,12 +35,12 @@ -- import qualified Language.Haskell.TH as TH import Capnp.Bits (WordCount)-import Capnp.Message (Mutability(..)) import qualified Capnp.Message as M+import Capnp.Mutability (Mutability(..), unsafeThaw) import Capnp.TraversalLimit (LimitT) import qualified Capnp.Untyped as U import Control.Monad.ST (RealWorld)--- import Internal.BuildPure (PureBuilder)+import Internal.BuildPure (PureBuilder) -- | Return a canonicalized message with a copy of the given struct as its -- root. returns a (message, segment) pair, where the segment is the first@@ -48,12 +49,13 @@ -- In addition to the usual reasons for failure when reading a message (traversal limit, -- malformed messages), this can fail if the message does not fit in a single segment, -- as the canonical form requires single-segment messages.-canonicalize- :: (U.RWCtx m s, M.MonadReadMessage mutIn m)- => U.Struct mutIn -> m (M.Message ('Mut s), M.Segment ('Mut s))+canonicalize :: U.RWCtx m s => U.Struct 'Const -> m (M.Message ('Mut s), M.Segment ('Mut s))+canonicalize s = unsafeThaw s >>= canonicalizeMut {-# SPECIALIZE canonicalize :: U.Struct 'Const -> LimitT IO (M.Message ('Mut RealWorld), M.Segment ('Mut RealWorld)) #-}-{-# SPECIALIZE canonicalize :: U.Struct ('Mut RealWorld) -> LimitT IO (M.Message ('Mut RealWorld), M.Segment ('Mut RealWorld)) #-}-canonicalize rootStructIn = do+{-# SPECIALIZE canonicalize :: U.Struct 'Const -> PureBuilder s (M.Message ('Mut s), M.Segment ('Mut s)) #-}++canonicalizeMut :: U.RWCtx m s => U.Struct ('Mut s) -> m (M.Message ('Mut s), M.Segment ('Mut s))+canonicalizeMut rootStructIn = do let msgIn = U.message @U.Struct rootStructIn -- Note [Allocation strategy] words <- totalWords msgIn@@ -62,6 +64,8 @@ U.setRoot rootStructOut segOut <- M.getSegment msgOut 0 pure (msgOut, segOut)+{-# SPECIALIZE canonicalizeMut :: U.Struct ('Mut RealWorld) -> LimitT IO (M.Message ('Mut RealWorld), M.Segment ('Mut RealWorld)) #-}+{-# SPECIALIZE canonicalizeMut :: U.Struct ('Mut s) -> PureBuilder s (M.Message ('Mut s), M.Segment ('Mut s)) #-} totalWords :: U.ReadCtx m mut => M.Message mut -> m WordCount totalWords msg = do@@ -72,22 +76,18 @@ M.numWords seg pure $ sum sizes -cloneCanonicalStruct- :: (U.RWCtx m s, M.MonadReadMessage mutIn m)- => U.Struct mutIn -> M.Message ('Mut s) -> m (U.Struct ('Mut s))-{-# SPECIALIZE cloneCanonicalStruct :: U.Struct 'Const -> M.Message ('Mut RealWorld) -> LimitT IO (U.Struct ('Mut RealWorld)) #-}+cloneCanonicalStruct :: U.RWCtx m s => U.Struct ('Mut s) -> M.Message ('Mut s) -> m (U.Struct ('Mut s)) {-# SPECIALIZE cloneCanonicalStruct :: U.Struct ('Mut RealWorld) -> M.Message ('Mut RealWorld) -> LimitT IO (U.Struct ('Mut RealWorld)) #-}+{-# SPECIALIZE cloneCanonicalStruct :: U.Struct ('Mut s) -> M.Message ('Mut s) -> PureBuilder s (U.Struct ('Mut s)) #-} cloneCanonicalStruct structIn msgOut = do (nWords, nPtrs) <- findCanonicalSectionCounts structIn structOut <- U.allocStruct msgOut (fromIntegral nWords) (fromIntegral nPtrs) copyCanonicalStruct structIn structOut pure structOut -copyCanonicalStruct- :: (U.RWCtx m s, M.MonadReadMessage mutIn m)- => U.Struct mutIn -> U.Struct ('Mut s) -> m ()-{-# SPECIALIZE copyCanonicalStruct :: U.Struct 'Const -> U.Struct ('Mut RealWorld) -> LimitT IO () #-}+copyCanonicalStruct :: U.RWCtx m s => U.Struct ('Mut s) -> U.Struct ('Mut s) -> m () {-# SPECIALIZE copyCanonicalStruct :: U.Struct ('Mut RealWorld) -> U.Struct ('Mut RealWorld) -> LimitT IO () #-}+{-# SPECIALIZE copyCanonicalStruct :: U.Struct ('Mut s) -> U.Struct ('Mut s) -> PureBuilder s () #-} copyCanonicalStruct structIn structOut = do let nWords = fromIntegral $ U.structWordCount structOut nPtrs = fromIntegral $ U.structPtrCount structOut@@ -100,8 +100,8 @@ U.setPtr ptrOut i structOut findCanonicalSectionCounts :: U.ReadCtx m mut => U.Struct mut -> m (Word16, Word16)-{-# SPECIALIZE findCanonicalSectionCounts :: U.Struct 'Const -> LimitT IO (Word16, Word16) #-} {-# SPECIALIZE findCanonicalSectionCounts :: U.Struct ('Mut RealWorld) -> LimitT IO (Word16, Word16) #-}+{-# SPECIALIZE findCanonicalSectionCounts :: U.Struct ('Mut s) -> PureBuilder s (Word16, Word16) #-} findCanonicalSectionCounts struct = do nWords <- canonicalSectionCount (== 0) (`U.getData` struct) (fromIntegral $ U.structWordCount struct) nPtrs <- canonicalSectionCount isNothing (`U.getPtr` struct) (fromIntegral $ U.structPtrCount struct)@@ -115,11 +115,9 @@ then canonicalSectionCount isDefault getIndex (total - 1) else pure $ fromIntegral total -cloneCanonicalPtr- :: (U.RWCtx m s, M.MonadReadMessage mutIn m)- => Maybe (U.Ptr mutIn) -> M.Message ('Mut s) -> m (Maybe (U.Ptr ('Mut s)))-{-# SPECIALIZE cloneCanonicalPtr :: Maybe (U.Ptr 'Const) -> M.Message ('Mut RealWorld) -> LimitT IO (Maybe (U.Ptr ('Mut RealWorld))) #-}+cloneCanonicalPtr :: U.RWCtx m s => Maybe (U.Ptr ('Mut s)) -> M.Message ('Mut s) -> m (Maybe (U.Ptr ('Mut s))) {-# SPECIALIZE cloneCanonicalPtr :: Maybe (U.Ptr ('Mut RealWorld)) -> M.Message ('Mut RealWorld) -> LimitT IO (Maybe (U.Ptr ('Mut RealWorld))) #-}+{-# SPECIALIZE cloneCanonicalPtr :: Maybe (U.Ptr ('Mut s)) -> M.Message ('Mut s) -> PureBuilder s (Maybe (U.Ptr ('Mut s))) #-} cloneCanonicalPtr ptrIn msgOut = case ptrIn of Nothing ->@@ -132,11 +130,9 @@ Just (U.PtrList list) -> Just . U.PtrList <$> cloneCanonicalList list msgOut -cloneCanonicalList- :: (U.RWCtx m s, M.MonadReadMessage mutIn m)- => U.List mutIn -> M.Message ('Mut s) -> m (U.List ('Mut s))-{-# SPECIALIZE cloneCanonicalList :: U.List 'Const -> M.Message ('Mut RealWorld) -> LimitT IO (U.List ('Mut RealWorld)) #-}+cloneCanonicalList :: U.RWCtx m s => U.List ('Mut s) -> M.Message ('Mut s) -> m (U.List ('Mut s)) {-# SPECIALIZE cloneCanonicalList :: U.List ('Mut RealWorld) -> M.Message ('Mut RealWorld) -> LimitT IO (U.List ('Mut RealWorld)) #-}+{-# SPECIALIZE cloneCanonicalList :: U.List ('Mut s) -> M.Message ('Mut s) -> PureBuilder s (U.List ('Mut s)) #-} cloneCanonicalList listIn msgOut = case listIn of U.List0 l -> U.List0 <$> U.allocList0 msgOut (U.length l)@@ -148,70 +144,24 @@ U.ListPtr l -> U.ListPtr <$> (U.allocListPtr msgOut (U.length l) >>= copyCanonicalPtrList l) U.ListStruct l -> U.ListStruct <$> cloneCanonicalStructList l msgOut -copyCanonicalDataList ::- ( U.RWCtx m s- , M.MonadReadMessage mutIn m- , U.ListItem r- , U.Unwrapped (U.Untyped r mutIn) ~ U.Unwrapped (U.Untyped r ('Mut s))- )- => U.ListOf r mutIn -> U.ListOf r ('Mut s) -> m (U.ListOf r ('Mut s))-{--{-# SPECIALIZE copyCanonicalDataList ::- ( U.ListItem r- , U.Unwrapped (U.Untyped r 'Const) ~ U.Unwrapped (U.Untyped r ('Mut RealWorld))- )- => U.ListOf r 'Const- -> U.ListOf r ('Mut RealWorld)- -> LimitT IO (U.ListOf r ('Mut RealWorld))- #-}--}-{-# SPECIALIZE copyCanonicalDataList ::- U.ListOf ('U.Data 'U.Sz8) 'Const- -> U.ListOf ('U.Data 'U.Sz8) ('Mut RealWorld)- -> LimitT IO (U.ListOf ('U.Data 'U.Sz8) ('Mut RealWorld))- #-}-{-# SPECIALIZE copyCanonicalDataList ::- U.ListOf ('U.Data 'U.Sz16) 'Const- -> U.ListOf ('U.Data 'U.Sz16) ('Mut RealWorld)- -> LimitT IO (U.ListOf ('U.Data 'U.Sz16) ('Mut RealWorld))- #-}-{-# SPECIALIZE copyCanonicalDataList ::- U.ListOf ('U.Data 'U.Sz32) 'Const- -> U.ListOf ('U.Data 'U.Sz32) ('Mut RealWorld)- -> LimitT IO (U.ListOf ('U.Data 'U.Sz32) ('Mut RealWorld))- #-}-{-# SPECIALIZE copyCanonicalDataList ::- U.ListOf ('U.Data 'U.Sz64) 'Const- -> U.ListOf ('U.Data 'U.Sz64) ('Mut RealWorld)- -> LimitT IO (U.ListOf ('U.Data 'U.Sz64) ('Mut RealWorld))- #-}-{-# SPECIALIZE copyCanonicalDataList ::- ( U.ListItem r- )- => U.ListOf r ('Mut RealWorld)- -> U.ListOf r ('Mut RealWorld)- -> LimitT IO (U.ListOf r ('Mut RealWorld))- #-}-copyCanonicalDataList listIn listOut = do- for_ [0..U.length listIn - 1] $ \i -> do- value <- U.index i listIn- U.setIndex value i listOut- pure listOut+copyCanonicalDataList lin lout = do+ U.copyListOf lout lin+ pure lout copyCanonicalPtrList- :: (U.RWCtx m s, M.MonadReadMessage mutIn m)- => U.ListOf ('U.Ptr 'Nothing) mutIn+ :: U.RWCtx m s+ => U.ListOf ('U.Ptr 'Nothing) ('Mut s) -> U.ListOf ('U.Ptr 'Nothing) ('Mut s) -> m (U.ListOf ('U.Ptr 'Nothing) ('Mut s)) {-# SPECIALIZE copyCanonicalPtrList- :: U.ListOf ('U.Ptr 'Nothing) 'Const+ :: U.ListOf ('U.Ptr 'Nothing) ('Mut RealWorld) -> U.ListOf ('U.Ptr 'Nothing) ('Mut RealWorld) -> LimitT IO (U.ListOf ('U.Ptr 'Nothing) ('Mut RealWorld)) #-} {-# SPECIALIZE copyCanonicalPtrList- :: U.ListOf ('U.Ptr 'Nothing) ('Mut RealWorld)- -> U.ListOf ('U.Ptr 'Nothing) ('Mut RealWorld)- -> LimitT IO (U.ListOf ('U.Ptr 'Nothing) ('Mut RealWorld))+ :: U.ListOf ('U.Ptr 'Nothing) ('Mut s)+ -> U.ListOf ('U.Ptr 'Nothing) ('Mut s)+ -> PureBuilder s (U.ListOf ('U.Ptr 'Nothing) ('Mut s)) #-} copyCanonicalPtrList listIn listOut = do for_ [0..U.length listIn - 1] $ \i -> do@@ -221,19 +171,19 @@ pure listOut cloneCanonicalStructList- :: (U.RWCtx m s, M.MonadReadMessage mutIn m)- => U.ListOf ('U.Ptr ('Just 'U.Struct)) mutIn+ :: U.RWCtx m s+ => U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut s) -> M.Message ('Mut s) -> m (U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut s)) {-# SPECIALIZE cloneCanonicalStructList- :: U.ListOf ('U.Ptr ('Just ('U.Struct))) 'Const+ :: U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut RealWorld) -> M.Message ('Mut RealWorld) -> LimitT IO (U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut RealWorld)) #-} {-# SPECIALIZE cloneCanonicalStructList- :: U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut RealWorld)- -> M.Message ('Mut RealWorld)- -> LimitT IO (U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut RealWorld))+ :: U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut s)+ -> M.Message ('Mut s)+ -> PureBuilder s (U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut s)) #-} cloneCanonicalStructList listIn msgOut = do (nWords, nPtrs) <- findCanonicalListSectionCounts listIn@@ -242,19 +192,19 @@ pure listOut copyCanonicalStructList- :: (U.RWCtx m s, M.MonadReadMessage mutIn m)- => U.ListOf ('U.Ptr ('Just 'U.Struct)) mutIn+ :: U.RWCtx m s+ => U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut s) -> U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut s) -> m () {-# SPECIALIZE copyCanonicalStructList- :: U.ListOf ('U.Ptr ('Just 'U.Struct)) 'Const+ :: U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut RealWorld) -> U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut RealWorld) -> LimitT IO () #-} {-# SPECIALIZE copyCanonicalStructList- :: U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut RealWorld)- -> U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut RealWorld)- -> LimitT IO ()+ :: U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut s)+ -> U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut s)+ -> PureBuilder s () #-} copyCanonicalStructList listIn listOut = for_ [0..U.length listIn - 1] $ \i -> do@@ -266,10 +216,10 @@ :: U.ReadCtx m mut => U.ListOf ('U.Ptr ('Just 'U.Struct)) mut -> m (Word16, Word16) {-# SPECIALIZE findCanonicalListSectionCounts- :: U.ListOf ('U.Ptr ('Just 'U.Struct)) 'Const -> LimitT IO (Word16, Word16)+ :: U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut RealWorld) -> LimitT IO (Word16, Word16) #-} {-# SPECIALIZE findCanonicalListSectionCounts- :: U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut RealWorld) -> LimitT IO (Word16, Word16)+ :: U.ListOf ('U.Ptr ('Just 'U.Struct)) ('Mut s) -> PureBuilder s (Word16, Word16) #-} findCanonicalListSectionCounts list = go 0 0 0 where go i !nWords !nPtrs
lib/Capnp/Message.hs view
@@ -37,15 +37,19 @@ , toByteString , fromByteString + -- * Accessing underlying storage+ , segToVecMut+ -- * Immutable messages , empty , singleSegment -- * Reading data from messages , MonadReadMessage(..)- , getSegment , getCap , getCapTable+ , getWord+ , totalNumWords -- * Mutable Messages , newMessage@@ -73,14 +77,13 @@ import Prelude hiding (read) -import Data.Bits (shiftL)- import Control.Monad (void, when, (>=>)) import Control.Monad.Catch (MonadThrow(..)) import Control.Monad.Primitive (PrimMonad, PrimState, stToPrim) import Control.Monad.State (evalStateT, get, put) import Control.Monad.Trans.Class (lift) import Control.Monad.Writer (execWriterT, tell)+import Data.Bits (shiftL) import Data.ByteString.Internal (ByteString(..)) import Data.Bytes.Get (getWord32le, runGetS) import Data.Maybe (fromJust)@@ -161,6 +164,15 @@ , used :: MutVar s WordCount } +-- | Return the underlying storage of a mutable segment, as a vector.+--+-- Note that the elements of the vector will be stored in little-endian form, regardless of+-- CPU endianness. This is a low level function that you should probably not use.+segToVecMut :: (PrimMonad m, PrimState m ~ s) => Segment ('Mut s) -> m (SMV.MVector s Word64)+segToVecMut (SegMut MutSegment{vec, used}) = do+ count <- readMutVar used+ pure $ SMV.take (fromIntegral count) vec+ instance Eq (MutSegment s) where MutSegment{used=x} == MutSegment{used=y} = x == y @@ -177,10 +189,9 @@ -- | 'numCaps' gets the number of capabilities in a message's capability -- table. numCaps :: Message mut -> m Int- -- | @'internalGetSeg' message index@ gets the segment at index 'index'- -- in 'message'. Most callers should use the 'getSegment' wrapper, instead- -- of calling this directly.- internalGetSeg :: Message mut -> Int -> m (Segment mut)+ -- | @'getSegment' message index@ gets the segment at index 'index'+ -- in 'message'.+ getSegment :: Message mut -> Int -> m (Segment mut) -- | @'internalGetCap' cap index@ reads a capability from the message's -- capability table, returning the client. does not check bounds. Callers -- should use getCap instead.@@ -204,13 +215,6 @@ toByteString (SegConst (ConstSegment vec)) = PS fptr offset len where (fptr, offset, len) = SV.unsafeToForeignPtr (SV.unsafeCast vec) --- | @'getSegment' message index@ fetches the given segment in the message.--- It throws a 'E.BoundsError' if the address is out of bounds.-getSegment :: (MonadThrow m, MonadReadMessage mut m) => Message mut -> Int -> m (Segment mut)-getSegment msg i = do- checkIndex i =<< numSegs msg- internalGetSeg msg i- -- | @'withCapTable'@ replaces the capability table in the message. withCapTable :: V.Vector Client -> Message 'Const -> Message 'Const withCapTable newCaps (MsgConst msg) = MsgConst $ msg { constCaps = newCaps }@@ -219,6 +223,10 @@ getCapTable :: Message 'Const -> V.Vector Client getCapTable (MsgConst ConstMsg{constCaps}) = constCaps +-- | 'getWord' gets the word referred to by the 'WordPtr'+getWord :: MonadReadMessage mut m => WordPtr mut -> m Word64+getWord WordPtr{pSegment, pAddr=WordAt{wordIndex}} = read pSegment wordIndex+ -- | @'getCap' message index@ gets the capability with the given index from -- the message. throws 'E.BoundsError' if the index is out -- of bounds.@@ -230,11 +238,11 @@ else msg `internalGetCap` i -- | @'setSegment' message index segment@ sets the segment at the given index--- in the message. It throws a 'E.BoundsError' if the address is out of bounds.+-- in the message. setSegment :: WriteCtx m s => Message ('Mut s) -> Int -> Segment ('Mut s) -> m ()-setSegment msg i seg = do- checkIndex i =<< numSegs msg- internalSetSeg msg i seg+setSegment (MsgMut MutMsg{mutSegs}) segIndex seg = do+ segs <- AppendVec.getVector <$> readMutVar mutSegs+ MV.write segs segIndex seg -- | @'setCap' message index cap@ sets the sets the capability at @index@ in -- the message's capability table to @cap@. If the index is out of bounds, a@@ -268,7 +276,7 @@ instance Monad m => MonadReadMessage 'Const m where numSegs (MsgConst ConstMsg{constSegs}) = pure $ V.length constSegs numCaps (MsgConst ConstMsg{constCaps}) = pure $ V.length constCaps- internalGetSeg (MsgConst ConstMsg{constSegs}) i = constSegs `V.indexM` i+ getSegment (MsgConst ConstMsg{constSegs}) i = constSegs `V.indexM` i internalGetCap (MsgConst ConstMsg{constCaps}) i = constCaps `V.indexM` i numWords (SegConst (ConstSegment vec)) = pure $ WordCount $ SV.length vec@@ -415,22 +423,13 @@ stToPrim $ GMV.length . AppendVec.getVector <$> readMutVar mutSegs numCaps (MsgMut MutMsg{mutCaps}) = stToPrim $ GMV.length . AppendVec.getVector <$> readMutVar mutCaps- internalGetSeg (MsgMut MutMsg{mutSegs}) i = stToPrim $ do+ getSegment (MsgMut MutMsg{mutSegs}) i = stToPrim $ do segs <- AppendVec.getVector <$> readMutVar mutSegs MV.read segs i internalGetCap (MsgMut MutMsg{mutCaps}) i = stToPrim $ do caps <- AppendVec.getVector <$> readMutVar mutCaps MV.read caps i ---- | @'internalSetSeg' message index segment@ sets the segment at the given--- index in the message. Most callers should use the 'setSegment' wrapper,--- instead of calling this directly.-internalSetSeg :: WriteCtx m s => Message ('Mut s) -> Int -> Segment ('Mut s) -> m ()-internalSetSeg (MsgMut MutMsg{mutSegs}) segIndex seg = do- segs <- AppendVec.getVector <$> readMutVar mutSegs- MV.write segs segIndex seg- -- | @'write' segment index value@ writes a value to the 64-bit word -- at the provided index. Consider using 'setWord' on the message, -- instead of calling this directly.@@ -503,13 +502,19 @@ -- Not enough space in the current segment; allocate a new one. -- the new segment's size should match the total size of existing segments -- but `maxSegmentSize` bounds how large it can get.- totalAllocation <- sum <$>- traverse (getSegment msg >=> numWords) [0..segIndex]+ totalAllocation <- totalNumWords msg ( newSegIndex, _ ) <- newSegment msg (min (max totalAllocation size) maxSegmentSize) -- This is guaranteed to succeed, since we just made a segment with -- at least size available space: fromJust <$> allocInSeg msg newSegIndex size +-- | Return the total number of words in the message, i.e. the sum of+-- the results of `numWords` on all segments.+totalNumWords :: MonadReadMessage mut m => Message mut -> m WordCount+totalNumWords msg = do+ lastSegIndex <- pred <$> numSegs msg+ sum <$> traverse (getSegment msg >=> numWords) [0..lastSegIndex]+ -- | 'empty' is an empty message, i.e. a minimal message with a null pointer as -- its root object. empty :: Message 'Const@@ -590,16 +595,16 @@ -> m (Message 'Const) freezeMsg freezeSeg freezeCaps msg@(MsgMut MutMsg{mutCaps}) = do len <- numSegs msg- constSegs <- V.generateM len (internalGetSeg msg >=> freezeSeg)+ constSegs <- V.generateM len (getSegment msg >=> freezeSeg) constCaps <- freezeCaps . AppendVec.getVector =<< readMutVar mutCaps pure $ MsgConst ConstMsg{constSegs, constCaps} -- | @'checkIndex' index length@ checkes that 'index' is in the range -- [0, length), throwing a 'BoundsError' if not.-checkIndex :: (Integral a, MonadThrow m) => a -> a -> m ()+checkIndex :: MonadThrow m => Int -> Int -> m () checkIndex i len = when (i < 0 || i >= len) $ throwM E.BoundsError- { E.index = fromIntegral i- , E.maxIndex = fromIntegral len+ { E.index = i+ , E.maxIndex = len }
lib/Capnp/Mutability.hs view
@@ -20,7 +20,7 @@ data Mutability = Const | Mut Type -- | 'MaybeMutable' relates mutable and immutable versions of a type.-class MaybeMutable (f :: Mutability -> *) where+class MaybeMutable (f :: Mutability -> Type) where -- | Convert an immutable value to a mutable one. thaw :: (PrimMonad m, PrimState m ~ s) => f 'Const -> m (f ('Mut s))
lib/Capnp/New/Rpc/Server.hs view
@@ -44,7 +44,7 @@ import Control.Exception.Safe (withException) import Control.Monad.STM.Class (MonadSTM(..)) import Data.Function ((&))-import Data.Kind (Constraint)+import Data.Kind (Constraint, Type) import qualified Data.Map.Strict as M import Data.Maybe (fromMaybe) import Data.Proxy (Proxy(..))@@ -92,7 +92,7 @@ -- if @'Server' i s@ is satisfied, @s@ is a server for interface @i@. -- The code generator generates a type class for each interface, and -- this will aways be an alias for that type class.- type Server i :: * -> Constraint+ type Server i :: Type -> Constraint -- | Convert the server to a 'MethodHandlerTree' populated with appropriate -- 'MethodHandler's for the interface. This is really only exported for use
lib/Capnp/Repr.hs view
@@ -71,9 +71,9 @@ import Prelude hiding (length) -import Capnp.Message (Mutability(..))-import qualified Capnp.Message as M-import Capnp.TraversalLimit (evalLimitT)+import qualified Capnp.Message as M+import Capnp.Mutability (MaybeMutable(..), Mutability(..))+import Capnp.TraversalLimit (evalLimitT) import Capnp.Untyped ( Allocate(..) , DataSz(..)@@ -83,6 +83,7 @@ , IsPtrRepr(..) , ListRepr(..) , ListReprFor+ , MaybePtr(..) , NormalListRepr(..) , PtrRepr(..) , Repr(..)@@ -92,14 +93,17 @@ , UntypedPtr , UntypedSomeList , UntypedSomePtr+ , Unwrapped )-import qualified Capnp.Untyped as U-import Data.Default (Default(..))+import qualified Capnp.Untyped as U+import Control.Monad.Primitive (PrimMonad, PrimState)+import Data.Default (Default(..)) import Data.Int-import Data.Kind (Type)-import Data.Maybe (fromJust)+import Data.Kind (Type)+import Data.Maybe (fromJust)+import Data.Traversable (for) import Data.Word-import GHC.Generics (Generic)+import GHC.Generics (Generic) -- | @'ReprFor' a@ denotes the Cap'n Proto wire represent of the type @a@. type family ReprFor (a :: Type) :: Repr@@ -186,6 +190,65 @@ instance U.MessageDefault (Raw a) => Default (Raw a 'Const) where def = fromJust $ evalLimitT maxBound $ U.messageDefault @(Raw a) M.empty +instance ReprMaybeMutable (ReprFor a) => MaybeMutable (Raw a) where+ thaw (Raw v) = Raw <$> rThaw @(ReprFor a) v+ freeze (Raw v) = Raw <$> rFreeze @(ReprFor a) v+ unsafeThaw (Raw v) = Raw <$> rUnsafeThaw @(ReprFor a) v+ unsafeFreeze (Raw v) = Raw <$> rUnsafeFreeze @(ReprFor a) v+ {-# INLINE thaw #-}+ {-# INLINE freeze #-}+ {-# INLINE unsafeThaw #-}+ {-# INLINE unsafeFreeze #-}++-- | Like MaybeMutable, but defined on the repr. Helper for implementing+-- MaybeMutable (Raw a)+class ReprMaybeMutable (r :: Repr) where+ rThaw :: (PrimMonad m, PrimState m ~ s) => Unwrapped (Untyped r 'Const) -> m (Unwrapped (Untyped r ('Mut s)))+ rUnsafeThaw :: (PrimMonad m, PrimState m ~ s) => Unwrapped (Untyped r 'Const) -> m (Unwrapped (Untyped r ('Mut s)))+ rFreeze :: (PrimMonad m, PrimState m ~ s) => Unwrapped (Untyped r ('Mut s)) -> m (Unwrapped (Untyped r 'Const))+ rUnsafeFreeze :: (PrimMonad m, PrimState m ~ s) => Unwrapped (Untyped r ('Mut s)) -> m (Unwrapped (Untyped r 'Const))++instance ReprMaybeMutable ('Ptr 'Nothing) where+ rThaw p = do+ MaybePtr p' <- thaw (MaybePtr p)+ pure p'+ rFreeze p = do+ MaybePtr p' <- freeze (MaybePtr p)+ pure p'+ rUnsafeThaw p = do+ MaybePtr p' <- unsafeThaw (MaybePtr p)+ pure p'+ rUnsafeFreeze p = do+ MaybePtr p' <- unsafeFreeze (MaybePtr p)+ pure p'++do+ let types =+ [ [t|'Just 'Struct|]+ , [t|'Just 'Cap|]+ , [t|'Just ('List 'Nothing)|]+ , [t|'Just ('List ('Just 'ListComposite))|]+ , [t|'Just ('List ('Just ('ListNormal 'NormalListPtr)))|]+ ]+ concat <$> for types (\t -> do+ [d|instance ReprMaybeMutable ('Ptr $t) where+ rThaw = thaw+ rFreeze = freeze+ rUnsafeThaw = thaw+ rUnsafeFreeze = freeze+ |])++instance ReprMaybeMutable ('Ptr ('Just ('List ('Just ('ListNormal ('NormalListData sz)))))) where+ rThaw = thaw+ rFreeze = freeze+ rUnsafeThaw = thaw+ rUnsafeFreeze = freeze++instance ReprMaybeMutable ('Data sz) where+ rThaw = pure+ rFreeze = pure+ rUnsafeThaw = pure+ rUnsafeFreeze = pure -- | Constraint that @a@ is a struct type. type IsStruct a = ReprFor a ~ 'Ptr ('Just 'Struct)
lib/Capnp/Repr/Methods.hs view
@@ -36,12 +36,13 @@ import qualified Capnp.New.Classes as NC import Capnp.New.Rpc.Common (Client(..), Pipeline(..)) import qualified Capnp.Repr as R-import Capnp.Rpc.Promise (newPromise)+import Capnp.Rpc.Promise (Promise, newPromise, wait) import qualified Capnp.Rpc.Server as Server import qualified Capnp.Rpc.Untyped as Rpc import Capnp.TraversalLimit (evalLimitT) import qualified Capnp.Untyped as U-import Control.Monad.Catch (MonadThrow)+import Control.Concurrent.STM (STM, atomically)+import Control.Monad.IO.Class (MonadIO(..)) import Control.Monad.STM.Class (MonadSTM(..)) import Data.Word import GHC.OverloadedLabels (IsLabel(..))@@ -85,30 +86,37 @@ -- | Call a method. Use the provided 'PureBuilder' to construct the parameters. callB- :: (AsClient f, R.IsCap c, R.IsStruct p, MonadSTM m)+ :: (AsClient f, R.IsCap c, R.IsStruct p, MonadIO m) => Method c p r -> (forall s. PureBuilder s (R.Raw p ('Mut s))) -> f c -> m (Pipeline r)-callB method buildRaw c = liftSTM $ do+callB method buildRaw c = liftIO $ do (params :: R.Raw a 'Const) <- R.Raw <$> createPure maxBound (R.fromRaw <$> buildRaw) callR method params c -- | Call a method, supplying the parameters as a 'Raw' struct. callR- :: (AsClient f, R.IsCap c, R.IsStruct p, MonadSTM m)+ :: (AsClient f, R.IsCap c, R.IsStruct p, MonadIO m) => Method c p r -> R.Raw p 'Const -> f c -> m (Pipeline r)-callR Method{interfaceId, methodId} (R.Raw arg) c = liftSTM $ do+callR method arg c = liftIO $ do+ p <- atomically (startCallR method arg c)+ Pipeline <$> wait p++startCallR+ :: (AsClient f, R.IsCap c, R.IsStruct p)+ => Method c p r -> R.Raw p 'Const -> f c -> STM (Promise Rpc.Pipeline)+startCallR Method{interfaceId, methodId} (R.Raw arg) c = do Client client <- asClient c (_, f) <- newPromise- Pipeline <$> Rpc.call+ Rpc.call Server.CallInfo { interfaceId , methodId , arguments = Just (U.PtrStruct arg) , response = f }- client+ client -- | Call a method, supplying the parmaeters in parsed form. callP@@ -117,11 +125,10 @@ , R.IsCap c , R.IsStruct p , NC.Parse p pp- , MonadSTM m- , MonadThrow m+ , MonadIO m ) => Method c p r -> pp -> f c -> m (Pipeline r)-callP method parsed client = do+callP method parsed client = liftIO $ do struct <- createPure maxBound $ do msg <- newMessage Nothing R.fromRaw <$> NC.encode msg parsed
lib/Capnp/Rpc/Transport.hs view
@@ -5,15 +5,20 @@ This module provides a 'Transport' type, which provides operations used to transmit messages between vats in the RPC protocol. -}-{-# LANGUAGE DataKinds #-}-{-# LANGUAGE TypeApplications #-}+{-# LANGUAGE DataKinds #-}+{-# LANGUAGE DuplicateRecordFields #-}+{-# LANGUAGE RecordWildCards #-}+{-# LANGUAGE TypeApplications #-} module Capnp.Rpc.Transport ( Transport(..) , handleTransport , socketTransport , tracingTransport+ , TraceConfig(..) ) where +import Prelude hiding (log)+ import Network.Socket (Socket) import System.IO (Handle) @@ -21,7 +26,9 @@ import Capnp.Convert (msgToParsed) import Capnp.IO (hGetMsg, hPutMsg, sGetMsg, sPutMsg) import Capnp.Message (Message, Mutability(Const))+import Capnp.New.Classes (Parsed) import Capnp.TraversalLimit (evalLimitT)+import Data.Default (def) import Text.Show.Pretty (ppShow) import qualified Capnp.Gen.Capnp.Rpc.New as R@@ -52,18 +59,39 @@ , recvMsg = sGetMsg socket limit } +data TraceConfig = TraceConfig+ { log :: String -> IO ()+ , showPayloads :: !Bool+ }+ -- | @'tracingTransport' log trans@ wraps another transport @trans@, loging -- messages when they are sent or received (using the @log@ function). This -- can be useful for debugging.-tracingTransport :: (String -> IO ()) -> Transport -> Transport-tracingTransport log trans = Transport+tracingTransport :: TraceConfig -> Transport -> Transport+tracingTransport tcfg trans = Transport { sendMsg = \msg -> do rpcMsg <- evalLimitT maxBound $ msgToParsed @R.Message msg- log $ "sending message: " ++ ppShow rpcMsg+ log tcfg $ "sending message: " ++ ppShow (editForTrace tcfg rpcMsg) sendMsg trans msg , recvMsg = do msg <- recvMsg trans rpcMsg <- evalLimitT maxBound $ msgToParsed @R.Message msg- log $ "received message: " ++ ppShow rpcMsg+ log tcfg $ "received message: " ++ ppShow (editForTrace tcfg rpcMsg) pure msg }++editForTrace :: TraceConfig -> Parsed R.Message -> Parsed R.Message+editForTrace tcfg rpcMsg =+ if showPayloads tcfg then+ rpcMsg+ else+ (case rpcMsg of+ R.Message (R.Message'call call) ->+ R.Message $ R.Message'call $+ call { R.params = def }+ R.Message (R.Message'return R.Return{union' = R.Return'results _, .. }) ->+ R.Message $ R.Message'return $+ R.Return { R.union' = R.Return'results def, .. }+ _ ->+ rpcMsg+ )
lib/Capnp/Rpc/Untyped.hs view
@@ -61,6 +61,7 @@ import Control.Monad.Trans.Class import Data.Word +import Capnp.Bits (WordCount, bytesToWordsFloor) import Capnp.New.Accessors import Control.Concurrent (threadDelay) import Control.Concurrent.Async (concurrently_, race_)@@ -109,7 +110,15 @@ , wrapException ) import Capnp.Rpc.Promise- (Fulfiller, breakOrFulfill, breakPromise, fulfill, newCallback)+ ( Fulfiller+ , Promise+ , breakOrFulfill+ , breakPromise+ , fulfill+ , newCallback+ , newPromise+ , newReadyPromise+ ) import Capnp.Rpc.Transport (Transport(recvMsg, sendMsg)) import Capnp.TraversalLimit (LimitT, defaultLimit, evalLimitT) import Internal.BuildPure (createPure)@@ -135,7 +144,7 @@ -- As an example, consider how we handle embargos: The 'Conn' type's 'embargos' -- table has values that are just 'Fulfiller's. This allows the code which triggers -- sending embargoes to have full control over what happens when they return,--- while the code that routes incoming messages (in 'coordinator') doesn't need+-- while the code that routes incoming messages (in 'recvLoop') doesn't need -- to concern itself with the details of embargos -- it just needs to route them -- to the right place. --@@ -207,32 +216,45 @@ | Dead data Conn' = Conn'- { sendQ :: TBQueue (Message 'Const)- , recvQ :: TBQueue (Message 'Const)- -- queues of messages to send and receive; each of these has a dedicated- -- thread doing the IO (see 'sendLoop' and 'recvLoop'):+ { sendQ :: TChan (Message 'Const, Fulfiller ())+ -- queue of messages to send sent to the remote vat; these are actually+ -- sent by a dedicated thread (see 'sendLoop').+ --+ -- The fulfiller is fulfilled after the message actually hits the transport.+ --+ -- The queue mainly exists for the sake of messages that are sent *while+ -- processing incomming messages*, since we cannot block in those cases,+ -- but it is used for all message sends to enforce ordering. The fulfiller+ -- is used by parts of the code (basically just calls) that want to block+ -- until their message is actually written to the socket. - , supervisor :: Supervisor+ , availableCallWords :: TVar WordCount+ -- Semaphore used to limit the memory that can be used by in-progress+ -- calls originating from this connection. We don't just use a TSem+ -- because waitTSem doesn't let us wait for more than one token with a+ -- single call.++ , supervisor :: Supervisor -- Supervisor managing the lifetimes of threads bound to this connection. - , questionIdPool :: IdPool- , exportIdPool :: IdPool+ , questionIdPool :: IdPool+ , exportIdPool :: IdPool -- Pools of identifiers for new questions and exports - , questions :: M.Map QAId EntryQA- , answers :: M.Map QAId EntryQA- , exports :: M.Map IEId EntryE- , imports :: M.Map IEId EntryI+ , questions :: M.Map QAId EntryQA+ , answers :: M.Map QAId EntryQA+ , exports :: M.Map IEId EntryE+ , imports :: M.Map IEId EntryI - , embargos :: M.Map EmbargoId (Fulfiller ())+ , embargos :: M.Map EmbargoId (Fulfiller ()) -- Outstanding embargos. When we receive a 'Disembargo' message with its -- context field set to receiverLoopback, we look up the embargo id in -- this table, and fulfill the promise. - , pendingCallbacks :: TQueue (IO ())+ , pendingCallbacks :: TQueue (IO ()) -- See Note [callbacks] - , bootstrap :: Maybe Client+ , bootstrap :: Maybe Client -- The capability which should be served as this connection's bootstrap -- interface (if any). }@@ -258,6 +280,14 @@ -- -- Defaults to 8192. + , maxCallWords :: !WordCount+ -- ^ The maximum total size of outstanding call messages that will be+ -- accepted; if this limit is reached, the implementation will not read+ -- more messages from the connection until some calls have completed+ -- and freed up enough space.+ --+ -- Defaults to 32MiB in words.+ , debugMode :: !Bool -- ^ In debug mode, errors reported by the RPC system to its peers will -- contain extra information. This should not be used in production, as@@ -289,6 +319,7 @@ def = ConnConfig { maxQuestions = 128 , maxExports = 8192+ , maxCallWords = bytesToWordsFloor $ 32 * 1024 * 1024 , debugMode = False , getBootstrap = \_ -> pure Nothing , withBootstrap = Nothing@@ -380,6 +411,7 @@ cfg@ConnConfig { maxQuestions , maxExports+ , maxCallWords , withBootstrap , debugMode }@@ -396,9 +428,10 @@ questionIdPool <- newIdPool maxQuestions exportIdPool <- newIdPool maxExports - sendQ <- newTBQueue $ fromIntegral maxQuestions- recvQ <- newTBQueue $ fromIntegral maxQuestions+ sendQ <- newTChan + availableCallWords <- newTVar maxCallWords+ questions <- M.new answers <- M.new exports <- M.new@@ -411,8 +444,8 @@ { supervisor = sup , questionIdPool , exportIdPool- , recvQ , sendQ+ , availableCallWords , questions , answers , exports@@ -430,9 +463,8 @@ pure (conn, conn') runConn (conn, conn') = do result <- try $- ( coordinator conn+ ( recvLoop transport conn `concurrently_` sendLoop transport conn'- `concurrently_` recvLoop transport conn' `concurrently_` callbacksLoop conn' ) `race_` useBootstrap conn conn'@@ -476,13 +508,6 @@ breakPromise fulfiller eDisconnected -- mark the connection as dead, making the live state inaccessible: writeTVar liveState Dead- -- Make sure any pending callbacks get run. This is important, since- -- some of these do things like raise disconnected exceptions.- --- -- FIXME: there's a race condition that we're not dealing with:- -- if the callbacks loop is killed between dequeuing an action and- -- performing it that action will be lost.- flushCallbacks conn' useBootstrap conn conn' = case withBootstrap of Nothing -> forever $ threadDelay maxBound@@ -512,8 +537,8 @@ -- | An entry in our questions or answers table. data EntryQA -- | An entry for which we have neither sent/received a finish, nor- -- a return. Contains two sets of callbacks, to invoke on each type- -- of message.+ -- a return. Contains two sets of callbacks, one to invoke on each+ -- type of message. = NewQA { onFinish :: SnocList (R.Parsed R.Finish -> STM ()) , onReturn :: SnocList (Return -> STM ())@@ -852,11 +877,11 @@ -- connection. -- | Queue a call on a client.-call :: MonadSTM m => Server.CallInfo -> Client -> m Pipeline+call :: MonadSTM m => Server.CallInfo -> Client -> m (Promise Pipeline) call Server.CallInfo { response } (Client Nothing) = liftSTM $ do breakPromise response eMethodUnimplemented state <- newTVar $ ReadyPipeline (Left eMethodUnimplemented)- pure Pipeline{state, steps = mempty}+ newReadyPromise Pipeline{state, steps = mempty} call info@Server.CallInfo { response } (Client (Just client')) = liftSTM $ do (localPipeline, response') <- makeLocalPipeline response let info' = info { Server.response = response' }@@ -867,7 +892,7 @@ q info' Nothing -> breakPromise response' eDisconnected- pure localPipeline+ newReadyPromise localPipeline PromiseClient { pState } -> readTVar pState >>= \case Ready { target } ->@@ -875,12 +900,12 @@ Embargo { callBuffer } -> do writeTQueue callBuffer info'- pure localPipeline+ newReadyPromise localPipeline Pending { tmpDest } -> case tmpDest of LocalDest LocalBuffer { callBuffer } -> do writeTQueue callBuffer info'- pure localPipeline+ newReadyPromise localPipeline RemoteDest AnswerDest { conn, answer } -> callRemote conn info $ AnswerTgt answer@@ -891,7 +916,7 @@ Error exn -> do breakPromise response' exn- pure localPipeline+ newReadyPromise localPipeline ImportClient cell -> do ImportRef { conn, importId } <- Fin.readCell cell@@ -913,7 +938,7 @@ pure (Pipeline{state, steps = mempty}, f') -- | Send a call to a remote capability.-callRemote :: Conn -> Server.CallInfo -> MsgTarget -> STM Pipeline+callRemote :: Conn -> Server.CallInfo -> MsgTarget -> STM (Promise Pipeline) callRemote conn Server.CallInfo{ interfaceId, methodId, arguments, response }@@ -921,13 +946,6 @@ conn'@Conn'{questions} <- getLive conn qid <- newQuestion conn' payload@Payload{capTable} <- makeOutgoingPayload conn arguments- sendCall conn' Call- { questionId = qid- , target = target- , params = payload- , interfaceId- , methodId- } -- save these in case the callee sends back releaseParamCaps = True in the return -- message: let paramCaps = catMaybes $ flip map (V.toList capTable) $ \R.CapDescriptor{union'} -> case union' of@@ -956,7 +974,18 @@ } qid questions- pure Pipeline { state = rp, steps = mempty }+ (p, f) <- newPromise+ f <- newCallback $ \r ->+ breakOrFulfill f (Pipeline { state = rp, steps = mempty } <$ r)+ sendCall conn' Call+ { questionId = qid+ , target = target+ , params = payload+ , interfaceId+ , methodId+ }+ f+ pure p -- | Callback to run when a return comes in that corresponds to a call -- we sent. Registered in callRemote. The first argument is a list of@@ -1132,66 +1161,75 @@ -- | See Note [callbacks] callbacksLoop :: Conn' -> IO ()-callbacksLoop Conn'{pendingCallbacks} = forever $ do- cbs <- atomically $ flushTQueue pendingCallbacks >>= \case- -- We need to make sure to block if there weren't any jobs, since- -- otherwise we'll busy loop, pegging the CPU.- [] -> retry- cbs -> pure cbs- sequence_ cbs---- Run the one iteration of the callbacks loop, without blocking.-flushCallbacks :: Conn' -> IO ()-flushCallbacks Conn'{pendingCallbacks} =- atomically (flushTQueue pendingCallbacks) >>= sequence_+callbacksLoop Conn'{pendingCallbacks} =+ loop `finally` cleanup+ where+ loop = forever $ doCallbacks $+ atomically $ flushTQueue pendingCallbacks >>= \case+ -- We need to make sure to block if there weren't any jobs, since+ -- otherwise we'll busy loop, pegging the CPU.+ [] -> retry+ cbs -> pure cbs+ cleanup =+ -- Make sure any pending callbacks get run. This is important, since+ -- some of these do things like raise disconnected exceptions.+ doCallbacks $ atomically $ flushTQueue pendingCallbacks+ doCallbacks getCbs =+ -- We need to be careful not to lose any callbacks in the event+ -- of an exception (even an async one):+ bracket+ getCbs+ (foldr finally (pure ()))+ (\_ -> pure ()) -- | 'sendLoop' shunts messages from the send queue into the transport. sendLoop :: Transport -> Conn' -> IO () sendLoop transport Conn'{sendQ} =- forever $ atomically (readTBQueue sendQ) >>= sendMsg transport---- | 'recvLoop' shunts messages from the transport into the receive queue.-recvLoop :: Transport -> Conn' -> IO ()-recvLoop transport Conn'{recvQ} =- forever $ recvMsg transport >>= atomically . writeTBQueue recvQ+ forever $ do+ (msg, f) <- atomically $ readTChan sendQ+ sendMsg transport msg+ atomically $ fulfill f () --- | The coordinator processes incoming messages.-coordinator :: Conn -> IO ()+-- | 'recvLoop' processes incoming messages.+recvLoop :: Transport -> Conn -> IO () -- The logic here mostly routes messages to other parts of the code that know -- more about the objects in question; See Note [Organization] for more info.-coordinator conn@Conn{debugMode} = forever $ atomically $ do- conn'@Conn'{recvQ} <- getLive conn- flip catchSTM (throwSTM . makeAbortExn debugMode) $ do- capnpMsg <- readTBQueue recvQ- evalLimitT defaultLimit $ do- rpcMsg <- msgToRaw capnpMsg- which <- structWhich rpcMsg- case which of- R.RW_Message'abort exn ->- parse exn >>= lift . handleAbortMsg conn- R.RW_Message'unimplemented oldMsg ->- parse oldMsg >>= lift . handleUnimplementedMsg conn- R.RW_Message'bootstrap bs ->- parse bs >>= lift . handleBootstrapMsg conn- R.RW_Message'call call ->- handleCallMsg conn call- R.RW_Message'return ret -> do- ret' <- acceptReturn conn ret- lift $ handleReturnMsg conn ret'- R.RW_Message'finish finish ->- parse finish >>= lift . handleFinishMsg conn- R.RW_Message'resolve res ->- parse res >>= lift . handleResolveMsg conn- R.RW_Message'release release ->- parse release >>= lift . handleReleaseMsg conn- R.RW_Message'disembargo disembargo ->- parse disembargo >>= lift . handleDisembargoMsg conn- _ -> do- msg <- parse rpcMsg- lift $ sendPureMsg conn' $ R.Message'unimplemented msg+recvLoop transport conn@Conn{debugMode} = forever $ do+ capnpMsg <- recvMsg transport+ atomically $ do+ flip catchSTM (throwSTM . makeAbortExn debugMode) $ do+ evalLimitT defaultLimit $ do+ rpcMsg <- msgToRaw capnpMsg+ which <- structWhich rpcMsg+ case which of+ R.RW_Message'abort exn ->+ parse exn >>= lift . handleAbortMsg conn+ R.RW_Message'unimplemented oldMsg ->+ parse oldMsg >>= lift . handleUnimplementedMsg conn+ R.RW_Message'bootstrap bs ->+ parse bs >>= lift . handleBootstrapMsg conn+ R.RW_Message'call call -> do+ handleCallMsg conn call+ R.RW_Message'return ret -> do+ ret' <- acceptReturn conn ret+ lift $ handleReturnMsg conn ret'+ R.RW_Message'finish finish ->+ parse finish >>= lift . handleFinishMsg conn+ R.RW_Message'resolve res ->+ parse res >>= lift . handleResolveMsg conn+ R.RW_Message'release release ->+ parse release >>= lift . handleReleaseMsg conn+ R.RW_Message'disembargo disembargo ->+ parse disembargo >>= lift . handleDisembargoMsg conn+ _ -> do+ msg <- parse rpcMsg+ lift $ do+ (_, onSent) <- newPromise+ conn' <- getLive conn+ sendPureMsg conn' (R.Message'unimplemented msg) onSent -- Each function handle*Msg handles a message of a particular type;--- 'coordinator' dispatches to these.+-- 'recvLoop' dispatches to these. handleAbortMsg :: Conn -> R.Parsed R.Exception -> STM () handleAbortMsg _ exn =@@ -1267,7 +1305,18 @@ handleCallMsg :: Conn -> Raw R.Call 'Const -> LimitT STM () handleCallMsg conn callMsg = do- conn'@Conn'{exports, answers} <- lift $ getLive conn+ conn'@Conn'{exports, answers, availableCallWords} <- lift $ getLive conn+ let capnpMsg = UntypedRaw.message @(Raw R.Call) callMsg++ -- Apply backpressure, by limiting the memory usage of outstanding call+ -- messages.+ msgWords <- Message.totalNumWords capnpMsg+ lift $ do+ available <- readTVar availableCallWords+ when (msgWords > available)+ retry+ writeTVar availableCallWords $! available - msgWords+ questionId <- parseField #questionId callMsg R.MessageTarget target <- parseField #target callMsg interfaceId <- parseField #interfaceId callMsg@@ -1283,7 +1332,10 @@ conn' (QAId questionId) NewQA- { onReturn = SnocList.empty+ { onReturn = SnocList.fromList+ [ \_ ->+ modifyTVar' availableCallWords (msgWords +)+ ] , onFinish = SnocList.fromList [ \R.Finish{releaseResultCaps} -> when releaseResultCaps $@@ -1381,12 +1433,15 @@ followPtrs (_:_) (Just _) = throwM $ eFailed "Tried to access pointer field of non-struct." -sendRawMsg :: Conn' -> Message 'Const -> STM ()-sendRawMsg conn' = writeTBQueue (sendQ conn')+sendRawMsg :: Conn' -> Message 'Const -> Fulfiller () -> STM ()+sendRawMsg conn' msg onSent = writeTChan (sendQ conn') (msg, onSent) -sendCall :: Conn' -> Call -> STM ()-sendCall conn' Call{questionId, target, interfaceId, methodId, params=Payload{content, capTable}} =- sendRawMsg conn' =<< createPure defaultLimit (do+sendCall :: Conn' -> Call -> Fulfiller () -> STM ()+sendCall+ conn'+ Call{questionId, target, interfaceId, methodId, params=Payload{content, capTable}}+ onSent = do+ msg <- createPure defaultLimit $ do mcontent <- traverse thaw content msg <- case mcontent of Just v -> pure $ UntypedRaw.message @UntypedRaw.Ptr v@@ -1403,65 +1458,71 @@ rpcMsg <- newRoot @R.Message () msg setVariant #call rpcMsg call pure msg- )+ sendRawMsg conn' msg onSent sendReturn :: Conn' -> Return -> STM ()-sendReturn conn' Return{answerId, releaseParamCaps, union'} = case union' of- Return'results Payload{content, capTable} ->- sendRawMsg conn' =<< createPure defaultLimit (do- mcontent <- traverse thaw content- msg <- case mcontent of- Just v -> pure $ UntypedRaw.message @UntypedRaw.Ptr v- Nothing -> Message.newMessage Nothing- payload <- new @R.Payload () msg- payload & setField #content (Raw mcontent)- payload & encodeField #capTable capTable- ret <- new @R.Return () msg- setVariant #results ret payload- ret & encodeField #answerId (qaWord answerId)- ret & encodeField #releaseParamCaps releaseParamCaps- rpcMsg <- newRoot @R.Message () msg- setVariant #return rpcMsg ret- pure msg- )- Return'exception exn ->- sendPureMsg conn' $ R.Message'return R.Return- { answerId = qaWord answerId- , releaseParamCaps- , union' = R.Return'exception exn- }- Return'canceled ->- sendPureMsg conn' $ R.Message'return R.Return- { answerId = qaWord answerId- , releaseParamCaps- , union' = R.Return'canceled- }- Return'resultsSentElsewhere ->- sendPureMsg conn' $ R.Message'return R.Return- { answerId = qaWord answerId- , releaseParamCaps- , union' = R.Return'resultsSentElsewhere- }- Return'takeFromOtherQuestion (QAId qid) ->- sendPureMsg conn' $ R.Message'return R.Return- { answerId = qaWord answerId- , releaseParamCaps- , union' = R.Return'takeFromOtherQuestion qid- }- Return'acceptFromThirdParty ptr ->- sendRawMsg conn' =<< createPure defaultLimit (do- mptr <- traverse thaw ptr- msg <- case mptr of- Just v -> pure $ UntypedRaw.message @UntypedRaw.Ptr v- Nothing -> Message.newMessage Nothing- ret <- new @R.Return () msg- ret & encodeField #answerId (qaWord answerId)- ret & encodeField #releaseParamCaps releaseParamCaps- setVariant #acceptFromThirdParty ret (Raw @(Maybe B.AnyPointer) mptr)- rpcMsg <- newRoot @R.Message () msg- setVariant #return rpcMsg ret- pure msg- )+sendReturn conn' Return{answerId, releaseParamCaps, union'} = do+ (_, onSent) <- newPromise+ case union' of+ Return'results Payload{content, capTable} -> do+ msg <- createPure defaultLimit $ do+ mcontent <- traverse thaw content+ msg <- case mcontent of+ Just v -> pure $ UntypedRaw.message @UntypedRaw.Ptr v+ Nothing -> Message.newMessage Nothing+ payload <- new @R.Payload () msg+ payload & setField #content (Raw mcontent)+ payload & encodeField #capTable capTable+ ret <- new @R.Return () msg+ setVariant #results ret payload+ ret & encodeField #answerId (qaWord answerId)+ ret & encodeField #releaseParamCaps releaseParamCaps+ rpcMsg <- newRoot @R.Message () msg+ setVariant #return rpcMsg ret+ pure msg+ sendRawMsg conn' msg onSent+ Return'exception exn ->+ sendPureMsg conn' (R.Message'return R.Return+ { answerId = qaWord answerId+ , releaseParamCaps+ , union' = R.Return'exception exn+ })+ onSent+ Return'canceled ->+ sendPureMsg conn' (R.Message'return R.Return+ { answerId = qaWord answerId+ , releaseParamCaps+ , union' = R.Return'canceled+ })+ onSent+ Return'resultsSentElsewhere ->+ sendPureMsg conn' (R.Message'return R.Return+ { answerId = qaWord answerId+ , releaseParamCaps+ , union' = R.Return'resultsSentElsewhere+ })+ onSent+ Return'takeFromOtherQuestion (QAId qid) ->+ sendPureMsg conn' (R.Message'return R.Return+ { answerId = qaWord answerId+ , releaseParamCaps+ , union' = R.Return'takeFromOtherQuestion qid+ })+ onSent+ Return'acceptFromThirdParty ptr -> do+ msg <- createPure defaultLimit $ do+ mptr <- traverse thaw ptr+ msg <- case mptr of+ Just v -> pure $ UntypedRaw.message @UntypedRaw.Ptr v+ Nothing -> Message.newMessage Nothing+ ret <- new @R.Return () msg+ ret & encodeField #answerId (qaWord answerId)+ ret & encodeField #releaseParamCaps releaseParamCaps+ setVariant #acceptFromThirdParty ret (Raw @(Maybe B.AnyPointer) mptr)+ rpcMsg <- newRoot @R.Message () msg+ setVariant #return rpcMsg ret+ pure msg+ sendRawMsg conn' msg onSent acceptReturn :: Conn -> Raw R.Return 'Const -> LimitT STM Return acceptReturn conn ret = do@@ -1502,13 +1563,15 @@ -- This can happen if we dropped the promise, but the release -- message is still in flight when the resolve message is sent. case union' of- R.Resolve'cap R.CapDescriptor{union' = R.CapDescriptor'receiverHosted importId} ->+ R.Resolve'cap R.CapDescriptor{union' = R.CapDescriptor'receiverHosted importId} -> do+ (_, onSent) <- newPromise -- Send a release message for the resolved cap, since -- we're not going to use it:- sendPureMsg conn' $ R.Message'release def+ sendPureMsg conn' (R.Message'release def { R.id = importId , R.referenceCount = 1- }+ })+ onSent -- Note [Level 3]: do we need to do something with -- thirdPartyHosted here? _ -> pure ()@@ -1620,13 +1683,15 @@ client <- Fin.readCell cell case client of ImportRef {conn=targetConn, importId}- | conn == targetConn ->- sendPureMsg conn' $ R.Message'disembargo R.Disembargo+ | conn == targetConn -> do+ (_, onSent) <- newPromise+ sendPureMsg conn' (R.Message'disembargo R.Disembargo { context = R.Disembargo'context' $ R.Disembargo'context'receiverLoopback embargoId , target = R.MessageTarget $ R.MessageTarget'importedCap (ieWord importId)- }+ })+ onSent _ -> abortDisembargoClient disembargoClient _ = abortDisembargoClient@@ -1644,8 +1709,11 @@ , info ] -- Note [Level 3]- go d conn' =- sendPureMsg conn' $ R.Message'unimplemented $ R.Message $ R.Message'disembargo d+ go d conn' = do+ (_, onSent) <- newPromise+ sendPureMsg conn'+ (R.Message'unimplemented $ R.Message $ R.Message'disembargo d)+ onSent lookupAbort :: (Eq k, Hashable k, Show k)@@ -1701,9 +1769,10 @@ capTable <- genSendableCapTableRaw conn content pure Payload { content, capTable } -sendPureMsg :: Conn' -> R.Parsed (Which R.Message) -> STM ()-sendPureMsg Conn'{sendQ} msg =- createPure maxBound (parsedToMsg (R.Message msg)) >>= writeTBQueue sendQ+sendPureMsg :: Conn' -> R.Parsed (Which R.Message) -> Fulfiller () -> STM ()+sendPureMsg Conn'{sendQ} msg onSent = do+ msg <- createPure maxBound (parsedToMsg (R.Message msg))+ writeTChan sendQ (msg, onSent) -- | Send a finish message, updating connection state and triggering -- callbacks as necessary.@@ -1713,7 +1782,8 @@ -- the return has also been received: subscribeReturn "question" conn questions (QAId questionId) $ \_ -> freeQuestion conn (QAId questionId)- sendPureMsg conn $ R.Message'finish finish+ (_, onSent) <- newPromise+ sendPureMsg conn (R.Message'finish finish) onSent updateQAFinish conn questions "question" finish -- | Send a return message, update the corresponding entry in our@@ -1795,11 +1865,11 @@ } val@HaveReturn{returnMsg} -> do- queueSTM conn (onRet returnMsg)+ onRet returnMsg pure val -- | Abort the connection, sending an abort message. This is only safe to call--- from within either the thread running the coordinator or the callback loop.+-- from within either the thread running the receieve loop or the callback loop. abortConn :: Conn' -> R.Parsed R.Exception -> STM a abortConn _ e = throwSTM (SentAbort e) @@ -1831,8 +1901,15 @@ } } pState <- newTVar Pending { tmpDest }- sendPureMsg conn' $- R.Message'bootstrap (def { R.questionId = qaWord qid } :: R.Parsed R.Bootstrap)++ -- Arguably, we should wait for this promise, since it's analagous+ -- to a call in terms of operation, but we only send one of these+ -- per connection, so whatever.+ (_, onSent) <- newPromise+ sendPureMsg conn'+ (R.Message'bootstrap (def { R.questionId = qaWord qid } :: R.Parsed R.Bootstrap))+ onSent+ M.insert NewQA { onReturn = SnocList.fromList@@ -1901,28 +1978,14 @@ disembargoAndResolve dest ( Just PromiseClient { origTarget=LocalDest _ }, RemoteDest dest) -> disembargoAndResolve dest- ( Nothing, RemoteDest dest ) ->- -- It's not clear to me what we should actually do if the promise- -- resolves to nullClient, but this can be encoded at the protocol- -- level, so we have to deal with it. Possible options:- --- -- 1. Perhaps this is simply illegal, and we should send an abort?- -- 2. Treat it as resolving to a local promise, in which case we- -- need to send a disembargo as above.- -- 3. Treat is as resolving to a remote promise, in which case we- -- can't send an embargo.- --- -- (3) doesn't seem possible to implement quite correctly, since- -- if we just resolve to nullClient right away, further calls will- -- start returning exceptions before outstanding calls return -- we- -- really do want to send a disembargo, but we can't because the- -- protocol insists that we don't if the promise resolves to a- -- remote cap.- --- -- What we currently do is (2); I(zenhack) intend to ask for- -- clarification on the mailing list.- disembargoAndResolve dest + ( Nothing, RemoteDest _ ) ->+ -- If it resolves to a null client, then we can't send a disembargo.+ -- Note that this may result in futrther calls throwing exceptions+ -- *before* the outstanding calls, which is a bit weird. But all+ -- calls with throw at some point, so it's probably fine.+ resolveNow+ -- Local promises never need embargos; we can just forward: ( _, LocalDest LocalBuffer { callBuffer } ) -> flushAndResolve callBuffer@@ -1995,11 +2058,13 @@ callback <- newCallback onEcho eid <- newEmbargo conn M.insert callback eid embargos- sendPureMsg conn $ R.Message'disembargo R.Disembargo+ (_, onSent) <- newPromise+ sendPureMsg conn (R.Message'disembargo R.Disembargo { target = marshalMsgTarget tgt , context = R.Disembargo'context' $ R.Disembargo'context'senderLoopback (embargoWord eid)- }+ })+ onSent -- | Resolve a promised client to the result of a return. See Note [resolveClient] --@@ -2226,12 +2291,14 @@ -- message with the correct count. releaseImport :: IEId -> Conn' -> STM () releaseImport importId conn'@Conn'{imports} = do+ (_, onSent) <- newPromise lookupAbort "imports" conn' imports importId $ \EntryI { remoteRc } ->- sendPureMsg conn' $ R.Message'release+ sendPureMsg conn' (R.Message'release R.Release { id = ieWord importId , referenceCount = remoteRc- }+ })+ onSent M.delete importId imports -- | Create a new client targeting an object in our answers table.
lib/Capnp/Untyped.hs view
@@ -75,7 +75,6 @@ , copyPtr , copyList , copyCap- , copyListOf , getClient , get, index , setIndex@@ -113,12 +112,15 @@ import Control.Monad.Primitive (PrimMonad(..)) import Control.Monad.ST (RealWorld) import Control.Monad.Trans.Class (MonadTrans(lift))+import Data.Coerce (coerce)+import Data.Function ((&)) import Data.Kind (Type) import qualified Data.ByteString as BS import qualified Language.Haskell.TH as TH -import Capnp.Address (OffsetError(..), WordAddr(..), pointerFrom)+import Capnp.Address+ (OffsetError(..), WordAddr(..), pointerFrom, resolveOffset) import Capnp.Bits ( BitCount(..) , ByteCount(..)@@ -131,10 +133,12 @@ ) import Capnp.Mutability (MaybeMutable(..), Mutability(..)) import Capnp.TraversalLimit (LimitT, MonadLimit(invoice))+import Internal.BuildPure (PureBuilder) -import qualified Capnp.Errors as E-import qualified Capnp.Message as M-import qualified Capnp.Pointer as P+import qualified Capnp.Errors as E+import qualified Capnp.Message as M+import qualified Capnp.Pointer as P+import qualified Data.Vector.Storable.Mutable as SMV ------------------------------------------------------------------------------- -- Untyped refernces to values in a message.@@ -176,7 +180,7 @@ -- | A list of values with representation 'r' in a message. newtype ListOf r mut = ListOf (ListRepOf r mut) -type family ListRepOf (r :: Repr) :: Mutability -> * where+type family ListRepOf (r :: Repr) :: Mutability -> Type where ListRepOf ('Ptr ('Just 'Struct)) = StructList ListRepOf r = NormalList @@ -197,6 +201,14 @@ checkListOf :: ReadCtx m mut => ListOf r mut -> m () + -- | Make a copy of the list, in the target message.+ copyListOf :: RWCtx m s => ListOf r ('Mut s) -> ListOf r ('Mut s) -> m ()+ {-# INLINE copyListOf #-}+ copyListOf dest src =+ forM_ [0..length src - 1] $ \i -> do+ value <- index i src+ setIndex value i dest+ default length :: (ListRepOf r ~ NormalList) => ListOf r mut -> Int length (ListOf nlist) = nLen nlist {-# INLINE length #-}@@ -240,6 +252,7 @@ (fromIntegral $ finiteBitSize (undefined :: Untyped r mut)) {-# INLINE checkListOf #-} + unsafeIndexBits :: forall a m mut. ( ReadCtx m mut@@ -312,6 +325,8 @@ {-# INLINE unsafeSetIndex #-} checkListOf _ = pure () {-# INLINE checkListOf #-}+ copyListOf _ _ = pure ()+ {-# INLINE copyListOf #-} instance ListItem ('Data 'Sz1) where unsafeIndex i (ListOf nlist) = do@@ -322,12 +337,21 @@ unsafeSetIndexBits @Word1 (Word1 value) i nlist {-# INLINE unsafeSetIndex #-} checkListOf (ListOf l) = checkNormalList l 1- {-# INLINE checkListOf #-}+ {-# INLINE copyListOf #-}+ copyListOf (ListOf dest) (ListOf src) = copyDataList dest src 1 -instance ListItem ('Data 'Sz8)-instance ListItem ('Data 'Sz16)-instance ListItem ('Data 'Sz32)-instance ListItem ('Data 'Sz64)+instance ListItem ('Data 'Sz8) where+ {-# INLINE copyListOf #-}+ copyListOf (ListOf dest) (ListOf src) = copyDataList dest src 8+instance ListItem ('Data 'Sz16) where+ {-# INLINE copyListOf #-}+ copyListOf (ListOf dest) (ListOf src) = copyDataList dest src 16+instance ListItem ('Data 'Sz32) where+ {-# INLINE copyListOf #-}+ copyListOf (ListOf dest) (ListOf src) = copyDataList dest src 32+instance ListItem ('Data 'Sz64) where+ {-# INLINE copyListOf #-}+ copyListOf (ListOf dest) (ListOf src) = copyDataList dest src 64 instance ListItem ('Ptr 'Nothing) where unsafeIndex i (ListOf (NormalList ptr@M.WordPtr{pAddr=addr@WordAt{..}} _)) =@@ -680,7 +704,7 @@ ------------------------------------------------------------------------------- -- | Types whose storage is owned by a message..-class HasMessage (f :: Mutability -> *) where+class HasMessage (f :: Mutability -> Type) where -- | Get the message in which the value is stored. message :: Unwrapped (f mut) -> M.Message mut @@ -764,110 +788,108 @@ get :: ReadCtx m mut => M.WordPtr mut -> m (Maybe (Ptr mut)) {-# INLINABLE get #-} {-# SPECIALIZE get :: M.WordPtr ('Mut RealWorld) -> LimitT IO (Maybe (Ptr ('Mut RealWorld))) #-}-get ptr@M.WordPtr{pMessage, pAddr} = do- word <- getWord ptr+{-# SPECIALIZE get :: M.WordPtr ('Mut s) -> PureBuilder s (Maybe (Ptr ('Mut s))) #-}+get ptr = do+ word <- M.getWord ptr case P.parsePtr word of- Nothing -> return Nothing- Just p -> case p of- P.CapPtr cap -> return $ Just $ PtrCap (CapAt pMessage cap)- P.StructPtr off dataSz ptrSz -> return $ Just $ PtrStruct $- StructAt ptr { M.pAddr = resolveOffset pAddr off } dataSz ptrSz- P.ListPtr off eltSpec -> Just <$>- getList ptr { M.pAddr = resolveOffset pAddr off } eltSpec- P.FarPtr twoWords offset segment -> do- landingSegment <- M.getSegment pMessage (fromIntegral segment)- let addr' = WordAt { wordIndex = fromIntegral offset- , segIndex = fromIntegral segment- }- let landingPtr = M.WordPtr+ Just (P.FarPtr twoWords offset segment) -> getFar ptr twoWords offset segment+ v -> getNear ptr v++getFar :: (M.MonadReadMessage mut m, MonadThrow m) => M.WordPtr mut -> Bool -> Word32 -> Word32 -> m (Maybe (Ptr mut))+getFar M.WordPtr{pMessage} twoWords offset segment = do+ landingSegment <- M.getSegment pMessage (fromIntegral segment)+ let addr' = WordAt { wordIndex = fromIntegral offset+ , segIndex = fromIntegral segment+ }+ let landingPtr = M.WordPtr+ { pMessage+ , pSegment = landingSegment+ , pAddr = addr'+ }+ landingPad <- M.getWord landingPtr+ if not twoWords+ then getNear landingPtr (P.parsePtr landingPad)+ else do+ case P.parsePtr landingPad of+ Just (P.FarPtr False off seg) -> do+ let segIndex = fromIntegral seg+ finalSegment <- M.getSegment pMessage segIndex+ tagWord <- M.getWord M.WordPtr { pMessage , pSegment = landingSegment- , pAddr = addr'+ , M.pAddr = addr' { wordIndex = wordIndex addr' + 1 } }- if not twoWords- then do- -- XXX: invoice so we don't open ourselves up to DoS- -- in the case of a chain of far pointers -- but a- -- better solution would be to just reject after the- -- first chain since this isn't actually legal. TODO- -- refactor (and then get rid of the MonadLimit- -- constraint).- invoice 1- get landingPtr- else do- landingPad <- getWord landingPtr- case P.parsePtr landingPad of- Just (P.FarPtr False off seg) -> do- let segIndex = fromIntegral seg- finalSegment <- M.getSegment pMessage segIndex- tagWord <- getWord M.WordPtr- { pMessage- , pSegment = landingSegment- , M.pAddr = addr' { wordIndex = wordIndex addr' + 1 }- }- let finalPtr = M.WordPtr- { pMessage- , pSegment = finalSegment- , pAddr = WordAt- { wordIndex = fromIntegral off- , segIndex- }- }- case P.parsePtr tagWord of- Just (P.StructPtr 0 dataSz ptrSz) ->- return $ Just $ PtrStruct $- StructAt finalPtr dataSz ptrSz- Just (P.ListPtr 0 eltSpec) ->- Just <$> getList finalPtr eltSpec- -- TODO: I'm not sure whether far pointers to caps are- -- legal; it's clear how they would work, but I don't- -- see a use, and the spec is unclear. Should check- -- how the reference implementation does this, copy- -- that, and submit a patch to the spec.- Just (P.CapPtr cap) ->- return $ Just $ PtrCap (CapAt pMessage cap)- ptr -> throwM $ E.InvalidDataError $- "The tag word of a far pointer's " ++- "2-word landing pad should be an intra " ++- "segment pointer with offset 0, but " ++- "we read " ++ show ptr- ptr -> throwM $ E.InvalidDataError $- "The first word of a far pointer's 2-word " ++- "landing pad should be another far pointer " ++- "(with a one-word landing pad), but we read " ++- show ptr+ let finalPtr = M.WordPtr+ { pMessage+ , pSegment = finalSegment+ , pAddr = WordAt+ { wordIndex = fromIntegral off+ , segIndex+ }+ }+ case P.parsePtr tagWord of+ Just (P.StructPtr 0 dataSz ptrSz) ->+ return $ Just $ PtrStruct $+ StructAt finalPtr dataSz ptrSz+ Just (P.ListPtr 0 eltSpec) ->+ Just . PtrList <$> getList finalPtr eltSpec+ -- TODO: I'm not sure whether far pointers to caps are+ -- legal; it's clear how they would work, but I don't+ -- see a use, and the spec is unclear. Should check+ -- how the reference implementation does this, copy+ -- that, and submit a patch to the spec.+ Just (P.CapPtr cap) ->+ return $ Just $ PtrCap (CapAt pMessage cap)+ ptr -> throwM $ E.InvalidDataError $+ "The tag word of a far pointer's " +++ "2-word landing pad should be an intra " +++ "segment pointer with offset 0, but " +++ "we read " ++ show ptr+ ptr -> throwM $ E.InvalidDataError $+ "The first word of a far pointer's 2-word " +++ "landing pad should be another far pointer " +++ "(with a one-word landing pad), but we read " +++ show ptr - where- getWord M.WordPtr{pSegment, pAddr=WordAt{wordIndex}} =- M.read pSegment wordIndex- resolveOffset addr@WordAt{..} off =- addr { wordIndex = wordIndex + fromIntegral off + 1 }- getList ptr@M.WordPtr{pAddr=addr@WordAt{wordIndex}} eltSpec = PtrList <$>- case eltSpec of- P.EltNormal sz len -> pure $ case sz of- P.Sz0 -> List0 (ListOf nlist)- P.Sz1 -> List1 (ListOf nlist)- P.Sz8 -> List8 (ListOf nlist)- P.Sz16 -> List16 (ListOf nlist)- P.Sz32 -> List32 (ListOf nlist)- P.Sz64 -> List64 (ListOf nlist)- P.SzPtr -> ListPtr (ListOf nlist)- where- nlist = NormalList ptr (fromIntegral len)- P.EltComposite _ -> do- tagWord <- getWord ptr- case P.parsePtr' tagWord of- P.StructPtr numElts dataSz ptrSz ->- pure $ ListStruct $ ListOf $ StructList- (StructAt- ptr { M.pAddr = addr { wordIndex = wordIndex + 1 } }- dataSz- ptrSz)- (fromIntegral numElts)- tag -> throwM $ E.InvalidDataError $- "Composite list tag was not a struct-" ++- "formatted word: " ++ show tag+getNear :: (M.MonadReadMessage mut m, MonadThrow m) => M.WordPtr mut -> Maybe P.Ptr -> m (Maybe (Ptr mut))+getNear ptr@M.WordPtr{pMessage, pAddr} = \case+ Nothing -> return Nothing+ Just p -> case p of+ P.CapPtr cap -> return $ Just $ PtrCap (CapAt pMessage cap)+ P.StructPtr off dataSz ptrSz -> return $ Just $ PtrStruct $+ StructAt ptr { M.pAddr = resolveOffset pAddr off } dataSz ptrSz+ P.ListPtr off eltSpec -> Just . PtrList <$>+ getList ptr { M.pAddr = resolveOffset pAddr off } eltSpec+ P.FarPtr{} -> throwM $ E.InvalidDataError+ "Unexpected far pointer where only near pointers were expected." +getList :: (M.MonadReadMessage mut m, MonadThrow m) => M.WordPtr mut -> P.EltSpec -> m (List mut)+getList ptr@M.WordPtr{pAddr=addr@WordAt{wordIndex}} eltSpec =+ case eltSpec of+ P.EltNormal sz len -> pure $ case sz of+ P.Sz0 -> List0 (ListOf nlist)+ P.Sz1 -> List1 (ListOf nlist)+ P.Sz8 -> List8 (ListOf nlist)+ P.Sz16 -> List16 (ListOf nlist)+ P.Sz32 -> List32 (ListOf nlist)+ P.Sz64 -> List64 (ListOf nlist)+ P.SzPtr -> ListPtr (ListOf nlist)+ where+ nlist = NormalList ptr (fromIntegral len)+ P.EltComposite _ -> do+ tagWord <- M.getWord ptr+ case P.parsePtr' tagWord of+ P.StructPtr numElts dataSz ptrSz ->+ pure $ ListStruct $ ListOf $ StructList+ (StructAt+ ptr { M.pAddr = addr { wordIndex = wordIndex + 1 } }+ dataSz+ ptrSz)+ (fromIntegral numElts)+ tag -> throwM $ E.InvalidDataError $+ "Composite list tag was not a struct-" +++ "formatted word: " ++ show tag+ -- | Return the EltSpec needed for a pointer to the given list. listEltSpec :: List msg -> P.EltSpec listEltSpec (ListStruct list@(ListOf (StructList (StructAt _ dataSz ptrSz) _))) =@@ -909,6 +931,9 @@ {-# SPECIALIZE setIndex :: ListItem r => Unwrapped (Untyped r ('Mut RealWorld)) -> Int -> ListOf r ('Mut RealWorld) -> LimitT IO () #-}+{-# SPECIALIZE setIndex+ :: ListItem r+ => Unwrapped (Untyped r ('Mut s)) -> Int -> ListOf r ('Mut s) -> PureBuilder s () #-} setIndex _ i list | i < 0 || length list <= i = throwM E.BoundsError { E.index = i, E.maxIndex = length list } setIndex value i list = unsafeSetIndex value i list@@ -920,6 +945,7 @@ setPointerTo :: M.WriteCtx m s => M.WordPtr ('Mut s) -> WordAddr -> P.Ptr -> m () {-# INLINABLE setPointerTo #-} {-# SPECIALIZE setPointerTo :: M.WordPtr ('Mut RealWorld) -> WordAddr -> P.Ptr -> LimitT IO () #-}+{-# SPECIALIZE setPointerTo :: M.WordPtr ('Mut s) -> WordAddr -> P.Ptr -> PureBuilder s () #-} setPointerTo M.WordPtr { pMessage = msg@@ -1000,6 +1026,7 @@ copyPtr :: RWCtx m s => M.Message ('Mut s) -> Maybe (Ptr ('Mut s)) -> m (Maybe (Ptr ('Mut s))) {-# INLINABLE copyPtr #-} {-# SPECIALIZE copyPtr :: M.Message ('Mut RealWorld) -> Maybe (Ptr ('Mut RealWorld)) -> LimitT IO (Maybe (Ptr ('Mut RealWorld))) #-}+{-# SPECIALIZE copyPtr :: M.Message ('Mut s) -> Maybe (Ptr ('Mut s)) -> PureBuilder s (Maybe (Ptr ('Mut s))) #-} copyPtr _ Nothing = pure Nothing copyPtr dest (Just (PtrCap cap)) = Just . PtrCap <$> copyCap dest cap copyPtr dest (Just (PtrList src)) = Just . PtrList <$> copyList dest src@@ -1015,6 +1042,7 @@ copyList :: RWCtx m s => M.Message ('Mut s) -> List ('Mut s) -> m (List ('Mut s)) {-# INLINABLE copyList #-} {-# SPECIALIZE copyList :: M.Message ('Mut RealWorld) -> List ('Mut RealWorld) -> LimitT IO (List ('Mut RealWorld)) #-}+{-# SPECIALIZE copyList :: M.Message ('Mut s) -> List ('Mut s) -> PureBuilder s (List ('Mut s)) #-} copyList dest src = case src of List0 src -> List0 <$> allocList0 dest (length src) List1 src -> List1 <$> copyNewListOf dest src allocList1@@ -1044,21 +1072,42 @@ copyListOf dest src pure dest +-- | @copyDataList dest src bits@ copies n elements of @src@ to @dest@, where n+-- is the length of the smaller list. @bits@ is the number of bits per element+-- in the two lists.+--+-- This should only used for non-pointer types, as it does not do a deep copy and+-- just copies the raw bytes.+--+-- Warning: if you get the @bits@ argument wrong, you may trample over data outside+-- the intended bounds.+copyDataList :: RWCtx m s => NormalList ('Mut s) -> NormalList ('Mut s) -> BitCount -> m ()+copyDataList dest src bits = do+ let unpack NormalList{nLen, nPtr = M.WordPtr{pSegment, pAddr=WordAt{wordIndex}}} =+ (nLen, wordIndex, pSegment) --- | Make a copy of the list, in the target message.-copyListOf- :: (ListItem r, RWCtx m s)- => ListOf r ('Mut s) -> ListOf r ('Mut s) -> m ()-{-# INLINE copyListOf #-}-copyListOf dest src =- forM_ [0..length src - 1] $ \i -> do- value <- index i src- setIndex value i dest+ (srcLen, srcOff, srcSeg) = unpack src+ (destLen, destOff, destSeg) = unpack dest + len = min destLen srcLen+ lenWords =+ fromIntegral len * bits+ & bitsToBytesCeil+ & bytesToWordsCeil++ sliceVec off =+ SMV.slice (fromIntegral off) (fromIntegral lenWords)+ srcVec <- M.segToVecMut srcSeg+ destVec <- M.segToVecMut destSeg+ SMV.copy+ (sliceVec destOff destVec)+ (sliceVec srcOff srcVec)+ -- | @'copyStruct' dest src@ copies the source struct to the destination struct. copyStruct :: RWCtx m s => Struct ('Mut s) -> Struct ('Mut s) -> m () {-# INLINABLE copyStruct #-} {-# SPECIALIZE copyStruct :: Struct ('Mut RealWorld) -> Struct ('Mut RealWorld) -> LimitT IO () #-}+{-# SPECIALIZE copyStruct :: Struct ('Mut s) -> Struct ('Mut s) -> PureBuilder s () #-} copyStruct dest src = do -- We copy both the data and pointer sections from src to dest, -- padding the tail of the destination section with zeros/null@@ -1084,6 +1133,8 @@ {-# INLINE index #-} {-# SPECIALIZE index :: ListItem r => Int -> ListOf r 'Const -> LimitT IO (Unwrapped (Untyped r 'Const)) #-} {-# SPECIALIZE index :: ListItem r => Int -> ListOf r ('Mut RealWorld) -> LimitT IO (Unwrapped (Untyped r ('Mut RealWorld))) #-}+{-# SPECIALIZE index :: ListItem r => Int -> ListOf r 'Const -> PureBuilder s (Unwrapped (Untyped r 'Const)) #-}+{-# SPECIALIZE index :: ListItem r => Int -> ListOf r ('Mut s) -> PureBuilder s (Unwrapped (Untyped r ('Mut s))) #-} index i list | i < 0 || i >= length list = throwM E.BoundsError { E.index = i, E.maxIndex = length list - 1 }@@ -1218,6 +1269,7 @@ invoicePtr :: MonadLimit m => Maybe (Ptr mut) -> m () {-# INLINABLE invoicePtr #-} {-# SPECIALIZE invoicePtr :: Maybe (Ptr ('Mut RealWorld)) -> LimitT IO () #-}+{-# SPECIALIZE invoicePtr :: Maybe (Ptr ('Mut s)) -> PureBuilder s () #-} invoicePtr p = invoice $! ptrInvoiceSize p ptrInvoiceSize :: Maybe (Ptr mut) -> WordCount@@ -1470,8 +1522,10 @@ concat <$> traverse mkWrappedInstance [ ''Ptr , ''List- , ''NormalList+ , ''Cap , ''Struct+ , ''NormalList+ , ''StructList ] do@@ -1516,3 +1570,13 @@ , "composite list" ) ]++instance MaybeMutable (IgnoreMut a) where+ thaw = pure . coerce+ freeze = pure . coerce++instance MaybeMutable MaybePtr where+ thaw (MaybePtr p) = MaybePtr <$> traverse thaw p+ freeze (MaybePtr p) = MaybePtr <$> traverse freeze p+ unsafeThaw (MaybePtr p) = MaybePtr <$> traverse unsafeThaw p+ unsafeFreeze (MaybePtr p) = MaybePtr <$> traverse unsafeFreeze p
lib/Internal/BuildPure.hs view
@@ -14,21 +14,20 @@ , createPure ) where -import Control.Monad.Catch (Exception, MonadThrow(..), SomeException)-import Control.Monad.Catch.Pure (CatchT, runCatchT)-import Control.Monad.Primitive (PrimMonad(..))-import Control.Monad.ST (ST)-import Control.Monad.Trans (MonadTrans(..))+import Control.Monad.Catch (Exception, MonadThrow(..), SomeException)+import Control.Monad.Primitive (PrimMonad(..))+import Control.Monad.ST (ST) import Capnp.Bits (WordCount) import Capnp.TraversalLimit (LimitT, MonadLimit, evalLimitT) import Capnp.Mutability+import Internal.STE -- | 'PureBuilder' is a monad transformer stack with the instnaces needed -- manipulate mutable messages. @'PureBuilder' s a@ is morally equivalent -- to @'LimitT' ('CatchT' ('ST' s)) a@-newtype PureBuilder s a = PureBuilder (LimitT (PrimCatchT (ST s)) a)+newtype PureBuilder s a = PureBuilder (LimitT (STE SomeException s) a) deriving(Functor, Applicative, Monad, MonadThrow, MonadLimit) instance PrimMonad (PureBuilder s) where@@ -36,7 +35,7 @@ primitive = PureBuilder . primitive runPureBuilder :: WordCount -> PureBuilder s a -> ST s (Either SomeException a)-runPureBuilder limit (PureBuilder m) = runPrimCatchT $ evalLimitT limit m+runPureBuilder limit (PureBuilder m) = steToST $ evalLimitT limit m -- | @'createPure' limit m@ creates a capnproto value in pure code according -- to @m@, then freezes it without copying. If @m@ calls 'throwM' then@@ -49,23 +48,3 @@ throwLeft :: (Exception e, MonadThrow m) => Either e a -> m a throwLeft (Left e) = throwM e throwLeft (Right a) = pure a---- | 'PrimCatchT' is a trivial wrapper around 'CatchT', which implements--- 'PrimMonad'. This is a temporary workaround for:------ https://github.com/ekmett/exceptions/issues/65------ If we can get that issue fixed, we can delete this and just bump the--- min bound on the exceptions package.-newtype PrimCatchT m a = PrimCatchT (CatchT m a)- deriving(Functor, Applicative, Monad, MonadThrow)--runPrimCatchT :: Monad m => PrimCatchT m a -> m (Either SomeException a)-runPrimCatchT (PrimCatchT m) = runCatchT m--instance MonadTrans PrimCatchT where- lift = PrimCatchT . lift--instance PrimMonad m => PrimMonad (PrimCatchT m) where- type PrimState (PrimCatchT m) = PrimState m- primitive = lift . primitive
+ lib/Internal/STE.hs view
@@ -0,0 +1,65 @@+{-# LANGUAGE DeriveDataTypeable #-}+{-# LANGUAGE DerivingStrategies #-}+{-# LANGUAGE FlexibleInstances #-}+{-# LANGUAGE GeneralizedNewtypeDeriving #-}+{-# LANGUAGE RankNTypes #-}+{-# LANGUAGE ScopedTypeVariables #-}+{-# LANGUAGE TypeFamilies #-}+-- | Simplified implementation of the monad-ste package,+-- with a few extras.+module Internal.STE+ ( STE+ , throwSTE+ , runSTE+ , liftST+ , steToST+ , steToIO+ ) where++import Control.Exception+import Control.Monad.Catch (MonadThrow(..))+import Control.Monad.Primitive (PrimMonad(..))+import Control.Monad.ST+import Control.Monad.ST.Unsafe+import Data.Typeable (Typeable)++newtype InternalErr e = InternalErr e+ deriving stock (Typeable)++instance Show (InternalErr e) where+ show _ = "(InternalErr _)"++instance Typeable e => Exception (InternalErr e)++newtype STE e s a = STE (IO a)+ deriving newtype (Functor, Applicative, Monad)++instance PrimMonad (STE e s) where+ type PrimState (STE e s) = s+ primitive = liftST . primitive++liftST :: ST s a -> STE e s a+liftST st = STE (unsafeSTToIO st)++throwSTE :: Exception e => e -> STE e s a+throwSTE e = STE (throwIO (InternalErr e))++runSTE :: Exception e => (forall s. STE e s a) -> Either e a+runSTE ste = runST (steToST ste)++steToST :: Typeable e => STE e s a -> ST s (Either e a)+steToST (STE io) = unsafeIOToST $ do+ res <- try io+ case res of+ Left (InternalErr e) -> pure $ Left e+ Right v -> pure $ Right v++steToIO :: forall e a. Exception e => STE e RealWorld a -> IO a+steToIO (STE io) = do+ res <- try io+ case res of+ Left (InternalErr (e :: e)) -> throwIO e+ Right v -> pure v++instance MonadThrow (STE SomeException s) where+ throwM = throwSTE . toException
tests/Module/Capnp/Canonicalize.hs view
@@ -58,7 +58,7 @@ ourImplCanonicalize struct = createPure maxBound $ do msg <- M.newMessage Nothing Raw rawStruct <- encode msg struct- (msg, _) <- canonicalize rawStruct+ (msg, _) <- canonicalizeMut rawStruct pure msg refImplCanonicalize :: Parsed B.AnyStruct -> IO (M.Message 'M.Const)
tests/WalkSchemaCodeGenRequest.hs view
@@ -54,7 +54,7 @@ bytes <- BS.readFile "tests/data/schema-codegenreq" root <- evalLimitT maxBound (bsToRaw bytes) endQuota <- execLimitT 4096 (reader root)- endQuota `shouldBe` 3374+ endQuota `shouldBe` 3409 where reader :: Raw Schema.CodeGeneratorRequest 'M.Const -> LimitT IO () reader req = do