[mle] fix MleRouter::HandleLinkAccept message validation (#4297)

This commit moves the message security check to the top of
MleRouter::HandleLinkAccept.
This commit is contained in:
Jonathan Hui
2019-11-06 07:02:19 -08:00
parent de8033fad4
commit 4aab013e06
4 changed files with 60 additions and 48 deletions
+2 -2
View File
@@ -2770,11 +2770,11 @@ void Mle::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageIn
break;
case Header::kCommandLinkAccept:
Get<MleRouter>().HandleLinkAccept(aMessage, aMessageInfo, keySequence);
Get<MleRouter>().HandleLinkAccept(aMessage, aMessageInfo, keySequence, neighbor);
break;
case Header::kCommandLinkAcceptAndRequest:
Get<MleRouter>().HandleLinkAcceptAndRequest(aMessage, aMessageInfo, keySequence);
Get<MleRouter>().HandleLinkAcceptAndRequest(aMessage, aMessageInfo, keySequence, neighbor);
break;
case Header::kCommandAdvertisement:
+42 -41
View File
@@ -768,21 +768,24 @@ exit:
otError MleRouter::HandleLinkAccept(const Message & aMessage,
const Ip6::MessageInfo &aMessageInfo,
uint32_t aKeySequence)
uint32_t aKeySequence,
Neighbor * aNeighbor)
{
return HandleLinkAccept(aMessage, aMessageInfo, aKeySequence, false);
return HandleLinkAccept(aMessage, aMessageInfo, aKeySequence, aNeighbor, false);
}
otError MleRouter::HandleLinkAcceptAndRequest(const Message & aMessage,
const Ip6::MessageInfo &aMessageInfo,
uint32_t aKeySequence)
uint32_t aKeySequence,
Neighbor * aNeighbor)
{
return HandleLinkAccept(aMessage, aMessageInfo, aKeySequence, true);
return HandleLinkAccept(aMessage, aMessageInfo, aKeySequence, aNeighbor, true);
}
otError MleRouter::HandleLinkAccept(const Message & aMessage,
const Ip6::MessageInfo &aMessageInfo,
uint32_t aKeySequence,
Neighbor * aNeighbor,
bool aRequest)
{
static const uint8_t dataRequestTlvs[] = {Tlv::kNetworkData};
@@ -790,7 +793,6 @@ otError MleRouter::HandleLinkAccept(const Message & aMessage,
otError error = OT_ERROR_NONE;
const otThreadLinkInfo *linkInfo = static_cast<const otThreadLinkInfo *>(aMessageInfo.GetLinkInfo());
Router * router;
Neighbor * neighbor;
Neighbor::State neighborState;
Mac::ExtAddress macAddr;
VersionTlv version;
@@ -803,10 +805,6 @@ otError MleRouter::HandleLinkAccept(const Message & aMessage,
RouteTlv route;
LeaderDataTlv leaderData;
LinkMarginTlv linkMargin;
ChallengeTlv challenge;
TlvRequestTlv tlvRequest;
aMessageInfo.GetPeerAddr().ToExtAddress(macAddr);
// Source Address
SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kSourceAddress, sizeof(sourceAddress), sourceAddress));
@@ -821,20 +819,45 @@ otError MleRouter::HandleLinkAccept(const Message & aMessage,
LogMleMessage("Receive Link Accept", aMessageInfo.GetPeerAddr(), sourceAddress.GetRloc16());
}
// Version
SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kVersion, sizeof(version), version));
VerifyOrExit(version.IsValid(), error = OT_ERROR_PARSE);
VerifyOrExit(IsActiveRouter(sourceAddress.GetRloc16()), error = OT_ERROR_PARSE);
routerId = GetRouterId(sourceAddress.GetRloc16());
router = mRouterTable.GetRouter(routerId);
neighborState = (router != NULL) ? router->GetState() : Neighbor::kStateInvalid;
// Response
SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kResponse, sizeof(response), response));
VerifyOrExit(response.IsValid(), error = OT_ERROR_PARSE);
// Remove stale neighbors
if ((neighbor = GetNeighbor(macAddr)) != NULL && neighbor->GetRloc16() != sourceAddress.GetRloc16())
// verify response
switch (neighborState)
{
RemoveNeighbor(*neighbor);
case Neighbor::kStateLinkRequest:
VerifyOrExit(memcmp(router->GetChallenge(), response.GetResponse(), router->GetChallengeSize()) == 0,
error = OT_ERROR_SECURITY);
break;
case Neighbor::kStateInvalid:
VerifyOrExit((mChallengeTimeout > 0) && (memcmp(mChallenge, response.GetResponse(), sizeof(mChallenge)) == 0),
error = OT_ERROR_SECURITY);
case Neighbor::kStateValid:
break;
default:
ExitNow(error = OT_ERROR_SECURITY);
}
// Remove stale neighbors
if (aNeighbor && aNeighbor->GetRloc16() != sourceAddress.GetRloc16())
{
RemoveNeighbor(*aNeighbor);
}
// Version
SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kVersion, sizeof(version), version));
VerifyOrExit(version.IsValid(), error = OT_ERROR_PARSE);
// Link-Layer Frame Counter
SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kLinkFrameCounter, sizeof(linkFrameCounter), linkFrameCounter));
VerifyOrExit(linkFrameCounter.IsValid(), error = OT_ERROR_PARSE);
@@ -863,32 +886,6 @@ otError MleRouter::HandleLinkAccept(const Message & aMessage,
linkMargin.SetLinkMargin(0);
}
VerifyOrExit(IsActiveRouter(sourceAddress.GetRloc16()), error = OT_ERROR_PARSE);
routerId = GetRouterId(sourceAddress.GetRloc16());
router = mRouterTable.GetRouter(routerId);
neighborState = (router != NULL) ? router->GetState() : Neighbor::kStateInvalid;
// verify response
switch (neighborState)
{
case Neighbor::kStateLinkRequest:
VerifyOrExit(memcmp(router->GetChallenge(), response.GetResponse(), router->GetChallengeSize()) == 0,
error = OT_ERROR_SECURITY);
break;
case Neighbor::kStateInvalid:
VerifyOrExit((mChallengeTimeout > 0) && (memcmp(mChallenge, response.GetResponse(), sizeof(mChallenge)) == 0),
error = OT_ERROR_SECURITY);
break;
case Neighbor::kStateValid:
break;
default:
ExitNow(error = OT_ERROR_INVALID_STATE);
}
switch (mRole)
{
case OT_DEVICE_ROLE_DISABLED:
@@ -968,6 +965,7 @@ otError MleRouter::HandleLinkAccept(const Message & aMessage,
}
// finish link synchronization
aMessageInfo.GetPeerAddr().ToExtAddress(macAddr);
router->SetExtAddress(macAddr);
router->SetRloc16(sourceAddress.GetRloc16());
router->SetLinkFrameCounter(linkFrameCounter.GetFrameCounter());
@@ -986,6 +984,9 @@ otError MleRouter::HandleLinkAccept(const Message & aMessage,
if (aRequest)
{
ChallengeTlv challenge;
TlvRequestTlv tlvRequest;
// Challenge
SuccessOrExit(error = Tlv::GetTlv(aMessage, Tlv::kChallenge, sizeof(challenge), challenge));
VerifyOrExit(challenge.IsValid(), error = OT_ERROR_PARSE);
+7 -2
View File
@@ -675,14 +675,19 @@ private:
void HandleDetachStart(void);
otError HandleChildStart(AttachMode aMode);
otError HandleLinkRequest(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo, Neighbor *aNeighbor);
otError HandleLinkAccept(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo, uint32_t aKeySequence);
otError HandleLinkAccept(const Message & aMessage,
const Ip6::MessageInfo &aMessageInfo,
uint32_t aKeySequence,
Neighbor * aNeighbor);
otError HandleLinkAccept(const Message & aMessage,
const Ip6::MessageInfo &aMessageInfo,
uint32_t aKeySequence,
Neighbor * aNeighbor,
bool aRequest);
otError HandleLinkAcceptAndRequest(const Message & aMessage,
const Ip6::MessageInfo &aMessageInfo,
uint32_t aKeySequence);
uint32_t aKeySequence,
Neighbor * aNeighbor);
otError HandleAdvertisement(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo);
otError HandleParentRequest(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo);
otError HandleChildIdRequest(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo, uint32_t aKeySequence);
+9 -3
View File
@@ -130,9 +130,15 @@ private:
void HandleDetachStart(void) {}
otError HandleChildStart(AttachMode) { return OT_ERROR_NONE; }
otError HandleLinkRequest(const Message &, const Ip6::MessageInfo &, Neighbor *) { return OT_ERROR_DROP; }
otError HandleLinkAccept(const Message &, const Ip6::MessageInfo &, uint32_t) { return OT_ERROR_DROP; }
otError HandleLinkAccept(const Message &, const Ip6::MessageInfo &, uint32_t, bool) { return OT_ERROR_DROP; }
otError HandleLinkAcceptAndRequest(const Message &, const Ip6::MessageInfo &, uint32_t) { return OT_ERROR_DROP; }
otError HandleLinkAccept(const Message &, const Ip6::MessageInfo &, uint32_t, Neighbor *) { return OT_ERROR_DROP; }
otError HandleLinkAccept(const Message &, const Ip6::MessageInfo &, uint32_t, Neighbor *, bool)
{
return OT_ERROR_DROP;
}
otError HandleLinkAcceptAndRequest(const Message &, const Ip6::MessageInfo &, uint32_t, Neighbor *)
{
return OT_ERROR_DROP;
}
otError HandleAdvertisement(const Message &, const Ip6::MessageInfo &) { return OT_ERROR_DROP; }
otError HandleParentRequest(const Message &, const Ip6::MessageInfo &) { return OT_ERROR_DROP; }
otError HandleChildIdRequest(const Message &, const Ip6::MessageInfo &, uint32_t) { return OT_ERROR_DROP; }