[secure-transport] add SetState() and log state changes (#9692)

This commit adds `SetState()` in `SecureTransport` class which
also logs the state changes. New helper methods `IsState{}()` is also
added to simplify the code.
This commit is contained in:
Abtin Keshavarzian
2023-12-07 08:50:08 -08:00
committed by GitHub
parent 5318d421cf
commit bf9944887d
2 changed files with 79 additions and 40 deletions
+65 -40
View File
@@ -138,17 +138,28 @@ void SecureTransport::FreeMbedtls(void)
mbedtls_ssl_free(&mSsl);
}
void SecureTransport::SetState(State aState)
{
VerifyOrExit(mState != aState);
LogInfo("State: %s -> %s", StateToString(mState), StateToString(aState));
mState = aState;
exit:
return;
}
Error SecureTransport::Open(ReceiveHandler aReceiveHandler, ConnectedHandler aConnectedHandler, void *aContext)
{
Error error;
VerifyOrExit(mState == kStateClosed, error = kErrorAlready);
VerifyOrExit(IsStateClosed(), error = kErrorAlready);
SuccessOrExit(error = mSocket.Open(&SecureTransport::HandleReceive, this));
mConnectedCallback.Set(aConnectedHandler, aContext);
mReceiveCallback.Set(aReceiveHandler, aContext);
mState = kStateOpen;
SetState(kStateOpen);
exit:
return error;
@@ -158,7 +169,7 @@ Error SecureTransport::Connect(const Ip6::SockAddr &aSockAddr)
{
Error error;
VerifyOrExit(mState == kStateOpen, error = kErrorInvalidState);
VerifyOrExit(IsStateOpen(), error = kErrorInvalidState);
mMessageInfo.SetPeerAddr(aSockAddr.GetAddress());
mMessageInfo.SetPeerPort(aSockAddr.mPort);
@@ -176,12 +187,10 @@ void SecureTransport::HandleReceive(void *aContext, otMessage *aMessage, const o
void SecureTransport::HandleReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo)
{
switch (mState)
{
case SecureTransport::kStateClosed:
ExitNow();
VerifyOrExit(!IsStateClosed());
case SecureTransport::kStateOpen:
if (IsStateOpen())
{
IgnoreError(mSocket.Connect(Ip6::SockAddr(aMessageInfo.GetPeerAddr(), aMessageInfo.GetPeerPort())));
mMessageInfo.SetPeerAddr(aMessageInfo.GetPeerAddr());
@@ -196,17 +205,16 @@ void SecureTransport::HandleReceive(Message &aMessage, const Ip6::MessageInfo &a
mMessageInfo.SetSockPort(aMessageInfo.GetSockPort());
SuccessOrExit(Setup(false));
break;
default:
}
else
{
// Once DTLS session is started, communicate only with a peer.
VerifyOrExit((mMessageInfo.GetPeerAddr() == aMessageInfo.GetPeerAddr()) &&
(mMessageInfo.GetPeerPort() == aMessageInfo.GetPeerPort()));
break;
}
#ifdef MBEDTLS_SSL_SRV_C
if (mState == SecureTransport::kStateConnecting)
if (IsStateConnecting())
{
IgnoreError(SetClientId(mMessageInfo.GetPeerAddr().mFields.m8, sizeof(mMessageInfo.GetPeerAddr().mFields)));
}
@@ -224,7 +232,7 @@ Error SecureTransport::Bind(uint16_t aPort)
{
Error error;
VerifyOrExit(mState == kStateOpen, error = kErrorInvalidState);
VerifyOrExit(IsStateOpen(), error = kErrorInvalidState);
VerifyOrExit(!mTransportCallback.IsSet(), error = kErrorAlready);
SuccessOrExit(error = mSocket.Bind(aPort, Ip6::kNetifUnspecified));
@@ -237,7 +245,7 @@ Error SecureTransport::Bind(TransportCallback aCallback, void *aContext)
{
Error error = kErrorNone;
VerifyOrExit(mState == kStateOpen, error = kErrorInvalidState);
VerifyOrExit(IsStateOpen(), error = kErrorInvalidState);
VerifyOrExit(!mSocket.IsBound(), error = kErrorAlready);
VerifyOrExit(!mTransportCallback.IsSet(), error = kErrorAlready);
@@ -252,9 +260,9 @@ Error SecureTransport::Setup(bool aClient)
int rval;
// do not handle new connection before guard time expired
VerifyOrExit(mState == kStateOpen, rval = MBEDTLS_ERR_SSL_TIMEOUT);
VerifyOrExit(IsStateOpen(), rval = MBEDTLS_ERR_SSL_TIMEOUT);
mState = kStateInitializing;
SetState(kStateInitializing);
mbedtls_ssl_init(&mSsl);
mbedtls_ssl_config_init(&mConf);
@@ -361,7 +369,6 @@ Error SecureTransport::Setup(bool aClient)
mReceiveMessage = nullptr;
mMessageSubType = Message::kSubTypeNone;
mState = kStateConnecting;
if (mCipherSuites[0] == MBEDTLS_TLS_ECJPAKE_WITH_AES_128_CCM_8)
{
@@ -374,14 +381,14 @@ Error SecureTransport::Setup(bool aClient)
}
#endif
mState = kStateConnecting;
SetState(kStateConnecting);
Process();
exit:
if ((mState == kStateInitializing) && (rval != 0))
if (IsStateInitializing() && (rval != 0))
{
mState = kStateOpen;
SetState(kStateOpen);
FreeMbedtls();
}
@@ -453,7 +460,7 @@ void SecureTransport::Close(void)
{
Disconnect();
mState = kStateClosed;
SetState(kStateClosed);
mTimerSet = false;
mTransportCallback.Clear();
@@ -463,10 +470,10 @@ void SecureTransport::Close(void)
void SecureTransport::Disconnect(void)
{
VerifyOrExit(mState == kStateConnecting || mState == kStateConnected);
VerifyOrExit(IsStateConnectingOrConnected());
mbedtls_ssl_close_notify(&mSsl);
mState = kStateCloseNotify;
SetState(kStateCloseNotify);
mTimer.Start(kGuardTimeNewConnectionMilli);
mMessageInfo.Clear();
@@ -561,7 +568,7 @@ Error SecureTransport::GetPeerCertificateBase64(unsigned char *aPeerCert, size_t
{
Error error = kErrorNone;
VerifyOrExit(mState == kStateConnected, error = kErrorInvalidState);
VerifyOrExit(IsStateConnected(), error = kErrorInvalidState);
#if (MBEDTLS_VERSION_NUMBER >= 0x03010000)
VerifyOrExit(mbedtls_base64_encode(aPeerCert, aCertBufferSize, aCertLength,
@@ -1004,22 +1011,15 @@ void SecureTransport::HandleTimer(Timer &aTimer)
void SecureTransport::HandleTimer(void)
{
switch (mState)
if (IsStateConnectingOrConnected())
{
case kStateConnecting:
case kStateConnected:
Process();
break;
case kStateCloseNotify:
mState = kStateOpen;
}
else if (IsStateCloseNotify())
{
SetState(kStateOpen);
mTimer.Stop();
mConnectedCallback.InvokeIfSet(false);
break;
default:
OT_ASSERT(false);
OT_UNREACHABLE_CODE(break);
}
}
@@ -1029,15 +1029,15 @@ void SecureTransport::Process(void)
bool shouldDisconnect = false;
int rval;
while ((mState == kStateConnecting) || (mState == kStateConnected))
while (IsStateConnectingOrConnected())
{
if (mState == kStateConnecting)
if (IsStateConnecting())
{
rval = mbedtls_ssl_handshake(&mSsl);
if (mSsl.MBEDTLS_PRIVATE(state) == MBEDTLS_SSL_HANDSHAKE_OVER)
{
mState = kStateConnected;
SetState(kStateConnected);
mConnectedCallback.InvokeIfSet(true);
}
}
@@ -1174,6 +1174,31 @@ exit:
return error;
}
#if OT_SHOULD_LOG_AT(OT_LOG_LEVEL_INFO)
const char *SecureTransport::StateToString(State aState)
{
static const char *const kStateStrings[] = {
"Closed", // (0) kStateClosed
"Open", // (1) kStateOpen
"Initializing", // (2) kStateInitializing
"Connecting", // (3) kStateConnecting
"Connected", // (4) kStateConnected
"CloseNotify", // (5) kStateCloseNotify
};
static_assert(0 == kStateClosed, "kStateClosed valid is incorrect");
static_assert(1 == kStateOpen, "kStateOpen valid is incorrect");
static_assert(2 == kStateInitializing, "kStateInitializing valid is incorrect");
static_assert(3 == kStateConnecting, "kStateConnecting valid is incorrect");
static_assert(4 == kStateConnected, "kStateConnected valid is incorrect");
static_assert(5 == kStateCloseNotify, "kStateCloseNotify valid is incorrect");
return kStateStrings[aState];
}
#endif
} // namespace MeshCoP
} // namespace ot
+14
View File
@@ -71,6 +71,7 @@
#include "common/callback.hpp"
#include "common/locator.hpp"
#include "common/log.hpp"
#include "common/message.hpp"
#include "common/random.hpp"
#include "common/timer.hpp"
@@ -467,6 +468,15 @@ private:
static constexpr size_t kSecureTransportKeyBlockSize = 40;
static constexpr size_t kSecureTransportRandomBufferSize = 32;
bool IsStateClosed(void) const { return mState == kStateClosed; }
bool IsStateOpen(void) const { return mState == kStateOpen; }
bool IsStateInitializing(void) const { return mState == kStateInitializing; }
bool IsStateConnecting(void) const { return mState == kStateConnecting; }
bool IsStateConnected(void) const { return mState == kStateConnected; }
bool IsStateCloseNotify(void) const { return mState == kStateCloseNotify; }
bool IsStateConnectingOrConnected(void) const { return mState == kStateConnecting || mState == kStateConnected; }
void SetState(State aState);
void FreeMbedtls(void);
Error Setup(bool aClient);
@@ -545,6 +555,10 @@ private:
void Process(void);
#if OT_SHOULD_LOG_AT(OT_LOG_LEVEL_INFO)
static const char *StateToString(State aState);
#endif
State mState;
int mCipherSuites[2];