[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.
This commit is contained in:
Abtin Keshavarzian
2024-11-19 08:56:21 -08:00
committed by GitHub
parent a2ea648b39
commit 013bb3eee5
2 changed files with 30 additions and 125 deletions
+28 -108
View File
@@ -205,18 +205,18 @@ void SecureTransport::HandleReceive(Message &aMessage, const Ip6::MessageInfo &a
#ifdef MBEDTLS_SSL_SRV_C #ifdef MBEDTLS_SSL_SRV_C
if (IsStateConnecting()) 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 #endif
Receive(aMessage); mReceiveMessage = &aMessage;
Process();
mReceiveMessage = nullptr;
exit: exit:
return; return;
} }
uint16_t SecureTransport::GetUdpPort(void) const { return mSocket.GetSockName().GetPort(); }
Error SecureTransport::Bind(uint16_t aPort) Error SecureTransport::Bind(uint16_t aPort)
{ {
Error error; Error error;
@@ -354,11 +354,11 @@ Error SecureTransport::Setup(bool aClient)
rval = mbedtls_ssl_setup(&mSsl, &mConf); rval = mbedtls_ssl_setup(&mSsl, &mConf);
VerifyOrExit(rval == 0); 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) if (mDatagramTransport)
{ {
mbedtls_ssl_set_timer_cb(&mSsl, this, &SecureTransport::HandleMbedtlsSetTimer, HandleMbedtlsGetTimer); mbedtls_ssl_set_timer_cb(&mSsl, this, HandleMbedtlsSetTimer, HandleMbedtlsGetTimer);
} }
if (mCipherSuite == kEcjpakeWithAes128Ccm8) if (mCipherSuite == kEcjpakeWithAes128Ccm8)
@@ -376,17 +376,6 @@ Error SecureTransport::Setup(bool aClient)
mReceiveMessage = nullptr; mReceiveMessage = nullptr;
mMessageSubType = Message::kSubTypeNone; 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); SetState(kStateConnecting);
Process(); Process();
@@ -456,8 +445,6 @@ void SecureTransport::Close(void)
mTimer.Stop(); mTimer.Stop();
} }
void SecureTransport::Disconnect(void) { Disconnect(kDisconnectedLocalClosed); }
void SecureTransport::Disconnect(ConnectEvent aEvent) void SecureTransport::Disconnect(ConnectEvent aEvent)
{ {
VerifyOrExit(IsStateConnectingOrConnected()); VerifyOrExit(IsStateConnectingOrConnected());
@@ -760,14 +747,6 @@ exit:
#endif // OPENTHREAD_CONFIG_TLS_API_ENABLE #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 SecureTransport::Send(Message &aMessage, uint16_t aLength)
{ {
Error error = kErrorNone; Error error = kErrorNone;
@@ -791,15 +770,6 @@ exit:
return error; return error;
} }
void SecureTransport::Receive(Message &aMessage)
{
mReceiveMessage = &aMessage;
Process();
mReceiveMessage = nullptr;
}
int SecureTransport::HandleMbedtlsTransmit(void *aContext, const unsigned char *aBuf, size_t aLength) int SecureTransport::HandleMbedtlsTransmit(void *aContext, const unsigned char *aBuf, size_t aLength)
{ {
return static_cast<SecureTransport *>(aContext)->HandleMbedtlsTransmit(aBuf, aLength); return static_cast<SecureTransport *>(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) int SecureTransport::HandleMbedtlsTransmit(const unsigned char *aBuf, size_t aLength)
{ {
Error error; Error error = kErrorNone;
int rval = 0; 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<uint16_t>(aLength)));
if (mTransportCallback.IsSet())
{ {
LogDebg("HandleMbedtlsTransmit DTLS"); error = mTransportCallback.Invoke(*message, mMessageInfo);
} }
#if OPENTHREAD_CONFIG_TLS_API_ENABLE
else else
{ {
LogDebg("HandleMbedtlsTransmit TLS"); error = mSocket.SendTo(*message, mMessageInfo);
} }
#endif
error = HandleSecureTransportSend(aBuf, static_cast<uint16_t>(aLength), mMessageSubType);
exit:
FreeMessageOnError(message, error);
mMessageSubType = Message::kSubTypeNone; mMessageSubType = Message::kSubTypeNone;
switch (error) switch (error)
@@ -836,7 +811,7 @@ int SecureTransport::HandleMbedtlsTransmit(const unsigned char *aBuf, size_t aLe
break; break;
default: default:
LogWarn("HandleMbedtlsTransmit: %s error", ErrorToString(error)); LogWarnOnError(error, "HandleMbedtlsTransmit");
rval = MBEDTLS_ERR_NET_SEND_FAILED; rval = MBEDTLS_ERR_NET_SEND_FAILED;
break; break;
} }
@@ -851,29 +826,16 @@ int SecureTransport::HandleMbedtlsReceive(void *aContext, unsigned char *aBuf, s
int SecureTransport::HandleMbedtlsReceive(unsigned char *aBuf, size_t aLength) int SecureTransport::HandleMbedtlsReceive(unsigned char *aBuf, size_t aLength)
{ {
int rval; int rval = MBEDTLS_ERR_SSL_WANT_READ;
uint16_t readLength;
if (mCipherSuite == kEcjpakeWithAes128Ccm8) VerifyOrExit(mReceiveMessage != nullptr);
{
LogDebg("HandleMbedtlsReceive DTLS");
}
#if OPENTHREAD_CONFIG_TLS_API_ENABLE
else
{
LogDebg("HandleMbedtlsReceive TLS");
}
#endif
VerifyOrExit(mReceiveMessage != nullptr && (rval = mReceiveMessage->GetLength() - mReceiveMessage->GetOffset()) > 0, readLength = mReceiveMessage->ReadBytes(mReceiveMessage->GetOffset(), aBuf, static_cast<uint16_t>(aLength));
rval = MBEDTLS_ERR_SSL_WANT_READ); VerifyOrExit(readLength > 0);
if (aLength > static_cast<size_t>(rval)) mReceiveMessage->MoveOffset(readLength);
{ rval = static_cast<int>(readLength);
aLength = static_cast<size_t>(rval);
}
rval = mReceiveMessage->ReadBytes(mReceiveMessage->GetOffset(), aBuf, static_cast<uint16_t>(aLength));
mReceiveMessage->MoveOffset(rval);
exit: exit:
return rval; return rval;
@@ -888,8 +850,6 @@ int SecureTransport::HandleMbedtlsGetTimer(void)
{ {
int rval; int rval;
LogDebg("HandleMbedtlsGetTimer");
if (!mTimerSet) if (!mTimerSet)
{ {
rval = -1; rval = -1;
@@ -917,17 +877,6 @@ void SecureTransport::HandleMbedtlsSetTimer(void *aContext, uint32_t aIntermedia
void SecureTransport::HandleMbedtlsSetTimer(uint32_t aIntermediate, uint32_t aFinish) 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) if (aFinish == 0)
{ {
mTimerSet = false; mTimerSet = false;
@@ -981,7 +930,6 @@ void SecureTransport::HandleMbedtlsExportKeys(mbedtls_ssl_key_export_type aType,
sha256.Update(keyBlock, kSecureTransportKeyBlockSize); sha256.Update(keyBlock, kSecureTransportKeyBlockSize);
sha256.Finish(kek); sha256.Finish(kek);
LogDebg("Generated KEK");
Get<KeyManager>().SetKek(kek.GetBytes()); Get<KeyManager>().SetKek(kek.GetBytes());
exit: exit:
@@ -1018,7 +966,6 @@ int SecureTransport::HandleMbedtlsExportKeys(const unsigned char *aMasterSecret,
sha256.Update(aKeyBlock, 2 * static_cast<uint16_t>(aMacLength + aKeyLength + aIvLength)); sha256.Update(aKeyBlock, 2 * static_cast<uint16_t>(aMacLength + aKeyLength + aIvLength));
sha256.Finish(kek); sha256.Finish(kek);
LogDebg("Generated KEK");
Get<KeyManager>().SetKek(kek.GetBytes()); Get<KeyManager>().SetKek(kek.GetBytes());
exit: exit:
@@ -1183,33 +1130,6 @@ void SecureTransport::HandleMbedtlsDebug(int aLevel, const char *aFile, int aLin
OT_UNUSED_VARIABLE(logLevel); 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) #if OT_SHOULD_LOG_AT(OT_LOG_LEVEL_INFO)
const char *SecureTransport::StateToString(State aState) const char *SecureTransport::StateToString(State aState)
+2 -17
View File
@@ -186,7 +186,7 @@ public:
* *
* @returns UDP port number. * @returns UDP port number.
*/ */
uint16_t GetUdpPort(void) const; uint16_t GetUdpPort(void) const { return mSocket.GetSockName().GetPort(); }
/** /**
* Binds with a transport callback. * Binds with a transport callback.
@@ -241,7 +241,7 @@ public:
/** /**
* Disconnects the session. * Disconnects the session.
*/ */
void Disconnect(void); void Disconnect(void) { Disconnect(kDisconnectedLocalClosed); }
/** /**
* Closes the socket. * Closes the socket.
@@ -413,18 +413,6 @@ public:
#endif // OPENTHREAD_CONFIG_TLS_API_ENABLE #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. * Sends data within the session.
* *
@@ -591,9 +579,6 @@ private:
static void HandleTimer(Timer &aTimer); static void HandleTimer(Timer &aTimer);
void HandleTimer(void); void HandleTimer(void);
Error HandleSecureTransportSend(const uint8_t *aBuf, uint16_t aLength, Message::SubType aMessageSubType);
void Receive(Message &aMessage);
void Process(void); void Process(void);
void Disconnect(ConnectEvent aEvent); void Disconnect(ConnectEvent aEvent);