module Network.TLS.Record.Recv (
recvRecord12,
recvRecord13,
) where
import qualified Data.ByteString as B
import Network.TLS.Context.Internal
import Network.TLS.Hooks
import Network.TLS.Imports
import Network.TLS.Packet
import Network.TLS.Record
import Network.TLS.Struct
import Network.TLS.Types
getMyPlainLimit :: Context -> IO Int
getMyPlainLimit :: Context -> IO Int
getMyPlainLimit Context
ctx = do
msiz <- Context -> IO (Maybe Int)
getMyRecordLimit Context
ctx
return $ case msiz of
Maybe Int
Nothing -> Int
defaultRecordSizeLimit
Just Int
siz -> Int
siz
getRecord
:: Context
-> Header
-> ByteString
-> IO (Either TLSError (Record Plaintext))
getRecord :: Context
-> Header -> ByteString -> IO (Either TLSError (Record Plaintext))
getRecord Context
ctx Header
header ByteString
content = do
Context -> (Logging -> IO ()) -> IO ()
withLog Context
ctx ((Logging -> IO ()) -> IO ()) -> (Logging -> IO ()) -> IO ()
forall a b. (a -> b) -> a -> b
$ \Logging
logging -> Logging -> Header -> ByteString -> IO ()
loggingIORecv Logging
logging Header
header ByteString
content
lim <- Context -> IO Int
getMyPlainLimit Context
ctx
runRxRecordState ctx $ do
let erecord = Header -> Fragment Ciphertext -> Record Ciphertext
forall a. Header -> Fragment a -> Record a
rawToRecord Header
header (Fragment Ciphertext -> Record Ciphertext)
-> Fragment Ciphertext -> Record Ciphertext
forall a b. (a -> b) -> a -> b
$ ByteString -> Fragment Ciphertext
fragmentCiphertext ByteString
content
decryptRecord erecord lim
exceedsTLSCiphertext :: Int -> Word16 -> Bool
exceedsTLSCiphertext :: Int -> Word16 -> Bool
exceedsTLSCiphertext Int
overhead Word16
actual =
Word16 -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral Word16
actual Int -> Int -> Bool
forall a. Ord a => a -> a -> Bool
> Int
defaultRecordSizeLimit Int -> Int -> Int
forall a. Num a => a -> a -> a
+ Int
overhead
recvRecord12
:: Context
-> IO (Either TLSError (Record Plaintext))
recvRecord12 :: Context -> IO (Either TLSError (Record Plaintext))
recvRecord12 Context
ctx =
Context -> Int -> IO (Either TLSError ByteString)
readExactBytes Context
ctx Int
5 IO (Either TLSError ByteString)
-> (Either TLSError ByteString
-> IO (Either TLSError (Record Plaintext)))
-> IO (Either TLSError (Record Plaintext))
forall a b. IO a -> (a -> IO b) -> IO b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= (TLSError -> IO (Either TLSError (Record Plaintext)))
-> (ByteString -> IO (Either TLSError (Record Plaintext)))
-> Either TLSError ByteString
-> IO (Either TLSError (Record Plaintext))
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext))
forall a. a -> IO a
forall (m :: * -> *) a. Monad m => a -> m a
return (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext)))
-> (TLSError -> Either TLSError (Record Plaintext))
-> TLSError
-> IO (Either TLSError (Record Plaintext))
forall b c a. (b -> c) -> (a -> b) -> a -> c
. TLSError -> Either TLSError (Record Plaintext)
forall a b. a -> Either a b
Left) (Either TLSError Header -> IO (Either TLSError (Record Plaintext))
recvLengthE (Either TLSError Header -> IO (Either TLSError (Record Plaintext)))
-> (ByteString -> Either TLSError Header)
-> ByteString
-> IO (Either TLSError (Record Plaintext))
forall b c a. (b -> c) -> (a -> b) -> a -> c
. ByteString -> Either TLSError Header
decodeHeader)
where
recvLengthE :: Either TLSError Header -> IO (Either TLSError (Record Plaintext))
recvLengthE = (TLSError -> IO (Either TLSError (Record Plaintext)))
-> (Header -> IO (Either TLSError (Record Plaintext)))
-> Either TLSError Header
-> IO (Either TLSError (Record Plaintext))
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext))
forall a. a -> IO a
forall (m :: * -> *) a. Monad m => a -> m a
return (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext)))
-> (TLSError -> Either TLSError (Record Plaintext))
-> TLSError
-> IO (Either TLSError (Record Plaintext))
forall b c a. (b -> c) -> (a -> b) -> a -> c
. TLSError -> Either TLSError (Record Plaintext)
forall a b. a -> Either a b
Left) Header -> IO (Either TLSError (Record Plaintext))
recvLength
recvLength :: Header -> IO (Either TLSError (Record Plaintext))
recvLength header :: Header
header@(Header ProtocolType
_ Version
_ Word16
readlen) = do
if Int -> Word16 -> Bool
exceedsTLSCiphertext Int
2048 Word16
readlen
then Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext))
forall a. a -> IO a
forall (m :: * -> *) a. Monad m => a -> m a
return (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext)))
-> Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext))
forall a b. (a -> b) -> a -> b
$ TLSError -> Either TLSError (Record Plaintext)
forall a b. a -> Either a b
Left TLSError
maximumSizeExceeded
else
Context -> Int -> IO (Either TLSError ByteString)
readExactBytes Context
ctx (Word16 -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral Word16
readlen)
IO (Either TLSError ByteString)
-> (Either TLSError ByteString
-> IO (Either TLSError (Record Plaintext)))
-> IO (Either TLSError (Record Plaintext))
forall a b. IO a -> (a -> IO b) -> IO b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= (TLSError -> IO (Either TLSError (Record Plaintext)))
-> (ByteString -> IO (Either TLSError (Record Plaintext)))
-> Either TLSError ByteString
-> IO (Either TLSError (Record Plaintext))
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext))
forall a. a -> IO a
forall (m :: * -> *) a. Monad m => a -> m a
return (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext)))
-> (TLSError -> Either TLSError (Record Plaintext))
-> TLSError
-> IO (Either TLSError (Record Plaintext))
forall b c a. (b -> c) -> (a -> b) -> a -> c
. TLSError -> Either TLSError (Record Plaintext)
forall a b. a -> Either a b
Left) (Context
-> Header -> ByteString -> IO (Either TLSError (Record Plaintext))
getRecord Context
ctx Header
header)
recvRecord13 :: Context -> IO (Either TLSError (Record Plaintext))
recvRecord13 :: Context -> IO (Either TLSError (Record Plaintext))
recvRecord13 Context
ctx = Context -> Int -> IO (Either TLSError ByteString)
readExactBytes Context
ctx Int
5 IO (Either TLSError ByteString)
-> (Either TLSError ByteString
-> IO (Either TLSError (Record Plaintext)))
-> IO (Either TLSError (Record Plaintext))
forall a b. IO a -> (a -> IO b) -> IO b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= (TLSError -> IO (Either TLSError (Record Plaintext)))
-> (ByteString -> IO (Either TLSError (Record Plaintext)))
-> Either TLSError ByteString
-> IO (Either TLSError (Record Plaintext))
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext))
forall a. a -> IO a
forall (m :: * -> *) a. Monad m => a -> m a
return (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext)))
-> (TLSError -> Either TLSError (Record Plaintext))
-> TLSError
-> IO (Either TLSError (Record Plaintext))
forall b c a. (b -> c) -> (a -> b) -> a -> c
. TLSError -> Either TLSError (Record Plaintext)
forall a b. a -> Either a b
Left) (Either TLSError Header -> IO (Either TLSError (Record Plaintext))
recvLengthE (Either TLSError Header -> IO (Either TLSError (Record Plaintext)))
-> (ByteString -> Either TLSError Header)
-> ByteString
-> IO (Either TLSError (Record Plaintext))
forall b c a. (b -> c) -> (a -> b) -> a -> c
. ByteString -> Either TLSError Header
decodeHeader)
where
recvLengthE :: Either TLSError Header -> IO (Either TLSError (Record Plaintext))
recvLengthE = (TLSError -> IO (Either TLSError (Record Plaintext)))
-> (Header -> IO (Either TLSError (Record Plaintext)))
-> Either TLSError Header
-> IO (Either TLSError (Record Plaintext))
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext))
forall a. a -> IO a
forall (m :: * -> *) a. Monad m => a -> m a
return (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext)))
-> (TLSError -> Either TLSError (Record Plaintext))
-> TLSError
-> IO (Either TLSError (Record Plaintext))
forall b c a. (b -> c) -> (a -> b) -> a -> c
. TLSError -> Either TLSError (Record Plaintext)
forall a b. a -> Either a b
Left) Header -> IO (Either TLSError (Record Plaintext))
recvLength
recvLength :: Header -> IO (Either TLSError (Record Plaintext))
recvLength header :: Header
header@(Header ProtocolType
_ Version
_ Word16
readlen) = do
if Int -> Word16 -> Bool
exceedsTLSCiphertext Int
256 Word16
readlen
then Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext))
forall a. a -> IO a
forall (m :: * -> *) a. Monad m => a -> m a
return (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext)))
-> Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext))
forall a b. (a -> b) -> a -> b
$ TLSError -> Either TLSError (Record Plaintext)
forall a b. a -> Either a b
Left TLSError
maximumSizeExceeded
else
Context -> Int -> IO (Either TLSError ByteString)
readExactBytes Context
ctx (Word16 -> Int
forall a b. (Integral a, Num b) => a -> b
fromIntegral Word16
readlen)
IO (Either TLSError ByteString)
-> (Either TLSError ByteString
-> IO (Either TLSError (Record Plaintext)))
-> IO (Either TLSError (Record Plaintext))
forall a b. IO a -> (a -> IO b) -> IO b
forall (m :: * -> *) a b. Monad m => m a -> (a -> m b) -> m b
>>= (TLSError -> IO (Either TLSError (Record Plaintext)))
-> (ByteString -> IO (Either TLSError (Record Plaintext)))
-> Either TLSError ByteString
-> IO (Either TLSError (Record Plaintext))
forall a c b. (a -> c) -> (b -> c) -> Either a b -> c
either (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext))
forall a. a -> IO a
forall (m :: * -> *) a. Monad m => a -> m a
return (Either TLSError (Record Plaintext)
-> IO (Either TLSError (Record Plaintext)))
-> (TLSError -> Either TLSError (Record Plaintext))
-> TLSError
-> IO (Either TLSError (Record Plaintext))
forall b c a. (b -> c) -> (a -> b) -> a -> c
. TLSError -> Either TLSError (Record Plaintext)
forall a b. a -> Either a b
Left) (Context
-> Header -> ByteString -> IO (Either TLSError (Record Plaintext))
getRecord Context
ctx Header
header)
maximumSizeExceeded :: TLSError
maximumSizeExceeded :: TLSError
maximumSizeExceeded = String -> AlertDescription -> TLSError
Error_Protocol String
"record exceeding maximum size" AlertDescription
RecordOverflow
readExactBytes :: Context -> Int -> IO (Either TLSError ByteString)
readExactBytes :: Context -> Int -> IO (Either TLSError ByteString)
readExactBytes Context
ctx Int
sz = do
hdrbs <- Context -> Int -> IO ByteString
contextRecv Context
ctx Int
sz
if B.length hdrbs == sz
then return $ Right hdrbs
else do
setEOF ctx
return . Left $
if B.null hdrbs
then Error_EOF
else
Error_Packet
( "partial packet: expecting "
++ show sz
++ " bytes, got: "
++ show (B.length hdrbs)
)