From 013bb3eee5e92f1e693fe7ef0fde93164ee88e55 Mon Sep 17 00:00:00 2001 From: Abtin Keshavarzian Date: Tue, 19 Nov 2024 08:56:21 -0800 Subject: [PATCH] [secure-transport] smaller enhancements in `SecureTransport` (#10932) This commit contains changes to `SecureTransport`: - Inlines simple methods (`Disconnect()` and `GetUdpPort()`). - Removes/refactors simpler methods, such as `Receive()` and `HandleSecureTransportSend()`, which are now directly moved into `HandleMbedtlsTransmit()` and `HandleReceive()`. - Removes unused `SetClientId()` and inlines its implementation. - Removes extra `LogDebug` line (can be tracked by logging `State` changes). - Simplifies `HandleMbedtlsReceive()` and how data is read from `mReceiveMessage`. - Uses a consistent style to refer to function pointers. --- src/core/meshcop/secure_transport.cpp | 136 ++++++-------------------- src/core/meshcop/secure_transport.hpp | 19 +--- 2 files changed, 30 insertions(+), 125 deletions(-) diff --git a/src/core/meshcop/secure_transport.cpp b/src/core/meshcop/secure_transport.cpp index 05652fd2e..f144a8d2a 100644 --- a/src/core/meshcop/secure_transport.cpp +++ b/src/core/meshcop/secure_transport.cpp @@ -205,18 +205,18 @@ void SecureTransport::HandleReceive(Message &aMessage, const Ip6::MessageInfo &a #ifdef MBEDTLS_SSL_SRV_C if (IsStateConnecting()) { - IgnoreError(SetClientId(mMessageInfo.GetPeerAddr().mFields.m8, sizeof(mMessageInfo.GetPeerAddr().mFields))); + mbedtls_ssl_set_client_transport_id(&mSsl, mMessageInfo.GetPeerAddr().GetBytes(), sizeof(Ip6::Address)); } #endif - Receive(aMessage); + mReceiveMessage = &aMessage; + Process(); + mReceiveMessage = nullptr; exit: return; } -uint16_t SecureTransport::GetUdpPort(void) const { return mSocket.GetSockName().GetPort(); } - Error SecureTransport::Bind(uint16_t aPort) { Error error; @@ -354,11 +354,11 @@ Error SecureTransport::Setup(bool aClient) rval = mbedtls_ssl_setup(&mSsl, &mConf); VerifyOrExit(rval == 0); - mbedtls_ssl_set_bio(&mSsl, this, &SecureTransport::HandleMbedtlsTransmit, HandleMbedtlsReceive, nullptr); + mbedtls_ssl_set_bio(&mSsl, this, HandleMbedtlsTransmit, HandleMbedtlsReceive, /* RecvTimeoutFn */ nullptr); if (mDatagramTransport) { - mbedtls_ssl_set_timer_cb(&mSsl, this, &SecureTransport::HandleMbedtlsSetTimer, HandleMbedtlsGetTimer); + mbedtls_ssl_set_timer_cb(&mSsl, this, HandleMbedtlsSetTimer, HandleMbedtlsGetTimer); } if (mCipherSuite == kEcjpakeWithAes128Ccm8) @@ -376,17 +376,6 @@ Error SecureTransport::Setup(bool aClient) mReceiveMessage = nullptr; mMessageSubType = Message::kSubTypeNone; - if (mCipherSuite == kEcjpakeWithAes128Ccm8) - { - LogInfo("DTLS started"); - } -#if OPENTHREAD_CONFIG_TLS_API_ENABLE - else - { - LogInfo("Application Secure (D)TLS started"); - } -#endif - SetState(kStateConnecting); Process(); @@ -456,8 +445,6 @@ void SecureTransport::Close(void) mTimer.Stop(); } -void SecureTransport::Disconnect(void) { Disconnect(kDisconnectedLocalClosed); } - void SecureTransport::Disconnect(ConnectEvent aEvent) { VerifyOrExit(IsStateConnectingOrConnected()); @@ -760,14 +747,6 @@ exit: #endif // OPENTHREAD_CONFIG_TLS_API_ENABLE -#ifdef MBEDTLS_SSL_SRV_C -Error SecureTransport::SetClientId(const uint8_t *aClientId, uint8_t aLength) -{ - int rval = mbedtls_ssl_set_client_transport_id(&mSsl, aClientId, aLength); - return Crypto::MbedTls::MapError(rval); -} -#endif - Error SecureTransport::Send(Message &aMessage, uint16_t aLength) { Error error = kErrorNone; @@ -791,15 +770,6 @@ exit: return error; } -void SecureTransport::Receive(Message &aMessage) -{ - mReceiveMessage = &aMessage; - - Process(); - - mReceiveMessage = nullptr; -} - int SecureTransport::HandleMbedtlsTransmit(void *aContext, const unsigned char *aBuf, size_t aLength) { return static_cast(aContext)->HandleMbedtlsTransmit(aBuf, aLength); @@ -807,22 +777,27 @@ int SecureTransport::HandleMbedtlsTransmit(void *aContext, const unsigned char * int SecureTransport::HandleMbedtlsTransmit(const unsigned char *aBuf, size_t aLength) { - Error error; - int rval = 0; + Error error = kErrorNone; + Message *message = mSocket.NewMessage(); + int rval; - if (mCipherSuite == kEcjpakeWithAes128Ccm8) + VerifyOrExit(message != nullptr, error = kErrorNoBufs); + message->SetSubType(mMessageSubType); + message->SetLinkSecurityEnabled(mLayerTwoSecurity); + + SuccessOrExit(error = message->AppendBytes(aBuf, static_cast(aLength))); + + if (mTransportCallback.IsSet()) { - LogDebg("HandleMbedtlsTransmit DTLS"); + error = mTransportCallback.Invoke(*message, mMessageInfo); } -#if OPENTHREAD_CONFIG_TLS_API_ENABLE else { - LogDebg("HandleMbedtlsTransmit TLS"); + error = mSocket.SendTo(*message, mMessageInfo); } -#endif - - error = HandleSecureTransportSend(aBuf, static_cast(aLength), mMessageSubType); +exit: + FreeMessageOnError(message, error); mMessageSubType = Message::kSubTypeNone; switch (error) @@ -836,7 +811,7 @@ int SecureTransport::HandleMbedtlsTransmit(const unsigned char *aBuf, size_t aLe break; default: - LogWarn("HandleMbedtlsTransmit: %s error", ErrorToString(error)); + LogWarnOnError(error, "HandleMbedtlsTransmit"); rval = MBEDTLS_ERR_NET_SEND_FAILED; break; } @@ -851,29 +826,16 @@ int SecureTransport::HandleMbedtlsReceive(void *aContext, unsigned char *aBuf, s int SecureTransport::HandleMbedtlsReceive(unsigned char *aBuf, size_t aLength) { - int rval; + int rval = MBEDTLS_ERR_SSL_WANT_READ; + uint16_t readLength; - if (mCipherSuite == kEcjpakeWithAes128Ccm8) - { - LogDebg("HandleMbedtlsReceive DTLS"); - } -#if OPENTHREAD_CONFIG_TLS_API_ENABLE - else - { - LogDebg("HandleMbedtlsReceive TLS"); - } -#endif + VerifyOrExit(mReceiveMessage != nullptr); - VerifyOrExit(mReceiveMessage != nullptr && (rval = mReceiveMessage->GetLength() - mReceiveMessage->GetOffset()) > 0, - rval = MBEDTLS_ERR_SSL_WANT_READ); + readLength = mReceiveMessage->ReadBytes(mReceiveMessage->GetOffset(), aBuf, static_cast(aLength)); + VerifyOrExit(readLength > 0); - if (aLength > static_cast(rval)) - { - aLength = static_cast(rval); - } - - rval = mReceiveMessage->ReadBytes(mReceiveMessage->GetOffset(), aBuf, static_cast(aLength)); - mReceiveMessage->MoveOffset(rval); + mReceiveMessage->MoveOffset(readLength); + rval = static_cast(readLength); exit: return rval; @@ -888,8 +850,6 @@ int SecureTransport::HandleMbedtlsGetTimer(void) { int rval; - LogDebg("HandleMbedtlsGetTimer"); - if (!mTimerSet) { rval = -1; @@ -917,17 +877,6 @@ void SecureTransport::HandleMbedtlsSetTimer(void *aContext, uint32_t aIntermedia void SecureTransport::HandleMbedtlsSetTimer(uint32_t aIntermediate, uint32_t aFinish) { - if (mCipherSuite == kEcjpakeWithAes128Ccm8) - { - LogDebg("SetTimer DTLS"); - } -#if OPENTHREAD_CONFIG_TLS_API_ENABLE - else - { - LogDebg("SetTimer TLS"); - } -#endif - if (aFinish == 0) { mTimerSet = false; @@ -981,7 +930,6 @@ void SecureTransport::HandleMbedtlsExportKeys(mbedtls_ssl_key_export_type aType, sha256.Update(keyBlock, kSecureTransportKeyBlockSize); sha256.Finish(kek); - LogDebg("Generated KEK"); Get().SetKek(kek.GetBytes()); exit: @@ -1018,7 +966,6 @@ int SecureTransport::HandleMbedtlsExportKeys(const unsigned char *aMasterSecret, sha256.Update(aKeyBlock, 2 * static_cast(aMacLength + aKeyLength + aIvLength)); sha256.Finish(kek); - LogDebg("Generated KEK"); Get().SetKek(kek.GetBytes()); exit: @@ -1183,33 +1130,6 @@ void SecureTransport::HandleMbedtlsDebug(int aLevel, const char *aFile, int aLin OT_UNUSED_VARIABLE(logLevel); } -Error SecureTransport::HandleSecureTransportSend(const uint8_t *aBuf, - uint16_t aLength, - Message::SubType aMessageSubType) -{ - Error error = kErrorNone; - Message *message = nullptr; - - VerifyOrExit((message = mSocket.NewMessage()) != nullptr, error = kErrorNoBufs); - message->SetSubType(aMessageSubType); - message->SetLinkSecurityEnabled(mLayerTwoSecurity); - - SuccessOrExit(error = message->AppendBytes(aBuf, aLength)); - - if (mTransportCallback.IsSet()) - { - SuccessOrExit(error = mTransportCallback.Invoke(*message, mMessageInfo)); - } - else - { - SuccessOrExit(error = mSocket.SendTo(*message, mMessageInfo)); - } - -exit: - FreeMessageOnError(message, error); - return error; -} - #if OT_SHOULD_LOG_AT(OT_LOG_LEVEL_INFO) const char *SecureTransport::StateToString(State aState) diff --git a/src/core/meshcop/secure_transport.hpp b/src/core/meshcop/secure_transport.hpp index 207842de8..b4b6e0240 100644 --- a/src/core/meshcop/secure_transport.hpp +++ b/src/core/meshcop/secure_transport.hpp @@ -186,7 +186,7 @@ public: * * @returns UDP port number. */ - uint16_t GetUdpPort(void) const; + uint16_t GetUdpPort(void) const { return mSocket.GetSockName().GetPort(); } /** * Binds with a transport callback. @@ -241,7 +241,7 @@ public: /** * Disconnects the session. */ - void Disconnect(void); + void Disconnect(void) { Disconnect(kDisconnectedLocalClosed); } /** * Closes the socket. @@ -413,18 +413,6 @@ public: #endif // OPENTHREAD_CONFIG_TLS_API_ENABLE -#ifdef MBEDTLS_SSL_SRV_C - /** - * Sets the Client ID used for generating the Hello Cookie. - * - * @param[in] aClientId A pointer to the Client ID. - * @param[in] aLength Number of bytes in the Client ID. - * - * @retval kErrorNone Successfully set the Client ID. - */ - Error SetClientId(const uint8_t *aClientId, uint8_t aLength); -#endif - /** * Sends data within the session. * @@ -591,9 +579,6 @@ private: static void HandleTimer(Timer &aTimer); void HandleTimer(void); - Error HandleSecureTransportSend(const uint8_t *aBuf, uint16_t aLength, Message::SubType aMessageSubType); - - void Receive(Message &aMessage); void Process(void); void Disconnect(ConnectEvent aEvent);