diff --git a/src/core/meshcop/border_agent.hpp b/src/core/meshcop/border_agent.hpp index 38cff7434..c99b6767b 100644 --- a/src/core/meshcop/border_agent.hpp +++ b/src/core/meshcop/border_agent.hpp @@ -341,7 +341,7 @@ private: void SendEnrollerResponse(Uri aUri, StateTlv::State aResponseState, const Coap::Message &aRequest); void SendEnrollerReportState(uint8_t aAdmitterState); Error AppendAdmitterTlvs(Coap::Message &aMessage, uint8_t aAdmitterState); - void ForwardUdpRelayToEnroller(const Coap::Message &aMessage); + void ForwardUdpRelayToEnroller(const Coap::Message &aMessage, bool aCheckEnrollerMode); void ForwardUdpProxyToEnroller(const Message &aMessage, const Ip6::MessageInfo &aMessageInfo); static Error ReadSteeringDataTlv(const Message &aMessage, SteeringData &aSteeringData); diff --git a/src/core/meshcop/border_agent_admitter.cpp b/src/core/meshcop/border_agent_admitter.cpp index 643b6e471..7e0b1ea7d 100644 --- a/src/core/meshcop/border_agent_admitter.cpp +++ b/src/core/meshcop/border_agent_admitter.cpp @@ -236,14 +236,16 @@ void Admitter::ForwardJoinerRelayToEnrollers(const Coap::Msg &aMsg) if (joiner != nullptr) { joiner->UpdateExpirationTime(); - iter.GetSessionAs()->ForwardUdpRelayToEnroller(aMsg.mMessage); + iter.GetSessionAs()->ForwardUdpRelayToEnroller(aMsg.mMessage, + /* aCheckEnrollerMode */ false); ExitNow(); } } for (EnrollerIterator iter(GetInstance()); !iter.IsDone(); iter.Advance()) { - iter.GetSessionAs()->ForwardUdpRelayToEnroller(aMsg.mMessage); + iter.GetSessionAs()->ForwardUdpRelayToEnroller(aMsg.mMessage, + /* aCheckEnrollerMode */ true); } exit: @@ -1425,12 +1427,16 @@ exit: return error; } -void Manager::CoapDtlsSession::ForwardUdpRelayToEnroller(const Coap::Message &aMessage) +void Manager::CoapDtlsSession::ForwardUdpRelayToEnroller(const Coap::Message &aMessage, bool aCheckEnrollerMode) { Error error = kErrorNone; VerifyOrExit(IsEnroller()); - VerifyOrExit(mEnroller->ShouldForwardJoinerRelay()); + + if (aCheckEnrollerMode) + { + VerifyOrExit(mEnroller->ShouldForwardJoinerRelay()); + } SuccessOrExit(error = ForwardUdpRelay(aMessage)); LogInfo("Forward %s to enroller - session %u", UriToString(), mIndex); diff --git a/tests/nexus/test_border_admitter.cpp b/tests/nexus/test_border_admitter.cpp index 56b7ff770..fb8206098 100644 --- a/tests/nexus/test_border_admitter.cpp +++ b/tests/nexus/test_border_admitter.cpp @@ -1326,8 +1326,10 @@ template bool DidFindAllEnrollers(const BitSetGet().Stop(); + // - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + Log("Send `EnrollerKeepAlive` message from `enrollers[0]` change its mode to disallow `kForwardJoinerRelayRx`"); + + modes[0] = MeshCoP::EnrollerModeTlv::kForwardUdpProxyRx; + + message = + enrollers[0]->Get().AllocateAndInitPriorityConfirmablePostMessage(kUriEnrollerKeepAlive); + VerifyOrQuit(message != nullptr); + + SuccessOrQuit(Tlv::Append(*message, MeshCoP::StateTlv::kAccept)); + SuccessOrQuit(Tlv::Append(*message, modes[0])); + + responseContexts[0].Clear(); + SuccessOrQuit(enrollers[0]->Get().SendMessage(*message, HandleResponse, &responseContexts[0])); + + nexus.AdvanceTime(Time::kOneSecondInMsec); + + VerifyOrQuit(responseContexts[0].mReceived); + VerifyOrQuit(responseContexts[0].mResponseState == MeshCoP::StateTlv::kAccept); + VerifyOrQuit(responseContexts[0].mHasAdmitterState); + + // - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + Log("Validate that the `enrollers[0]` mode got changed on `admitter`"); + + foundEnrollers.Clear(); + iter.Init(admitter.GetInstance()); + + while (iter.GetNextEnrollerInfo(enrollerInfo) == kErrorNone) + { + uint8_t matchedIndex; + + LogEnroller(enrollerInfo); + + matchedIndex = FindMatchingEnroller(enrollerInfo, kEnrollerIds, foundEnrollers); + + VerifyOrQuit(AsCoreType(&enrollerInfo.mSteeringData) == steeringData); + VerifyOrQuit(enrollerInfo.mMode == modes[matchedIndex]); + + if (matchedIndex == 0) + { + SuccessOrQuit(iter.GetNextJoinerInfo(joinerInfo)); + VerifyOrQuit(AsCoreType(&joinerInfo.mIid) == joinerIids[0]); + LogJoiner(joinerInfo); + } + + VerifyOrQuit(iter.GetNextJoinerInfo(joinerInfo) == kErrorNotFound); + } + + VerifyOrQuit(DidFindAllEnrollers(foundEnrollers)); + + nexus.AdvanceTime(Time::kOneSecondInMsec); + + for (uint8_t i = 0; i < kNumEnrollers; i++) + { + recvContext[i].Clear(); + } + + // - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + Log("Start `joiners[0]` again and validate that its `RelayRx` is still only forwarded to `enrollers[0]`"); + + joiners[0]->Get().Up(); + SuccessOrQuit(joiners[0]->Get().Start(kPskd, + /* aProvisioningUrl */ nullptr, + /* aVendorName */ nullptr, + /* aVendorModel */ nullptr, + /* aVendorSwVersion */ nullptr, + /* aVendorData */ nullptr, + /* aCallback */ nullptr, + /* aContext */ nullptr)); + + nexus.AdvanceTime(8 * Time::kOneSecondInMsec); + + for (uint8_t i = 0; i < kNumEnrollers; i++) + { + Ip6::InterfaceIdentifier readIid; + uint16_t joinerRouterRloc; + + message = AsCoapMessagePtr(recvContext[i].mRelayRxMsgs.GetHead()); + + if (i != 0) + { + VerifyOrQuit(message == nullptr); + continue; + } + + VerifyOrQuit(message != nullptr); + + VerifyOrQuit(message->ReadType() == Coap::kTypeNonConfirmable); + VerifyOrQuit(message->ReadCode() == Coap::kCodePost); + SuccessOrQuit(Tlv::Find(*message, readIid)); + SuccessOrQuit(Tlv::Find(*message, joinerRouterRloc)); + + VerifyOrQuit(readIid == joinerIids[0]); + VerifyOrQuit(joinerRouterRloc == admitter.Get().GetRloc16()); + } + + joiners[0]->Get().Stop(); + + for (uint8_t i = 0; i < kNumEnrollers; i++) + { + recvContext[i].Clear(); + } + + // - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + Log("Start `joiners[1]` and validate that its `RelayRx` is not longer forwarded to `enrollers[0]`"); + + joiners[1]->Get().Up(); + SuccessOrQuit(joiners[1]->Get().Start(kPskd, + /* aProvisioningUrl */ nullptr, + /* aVendorName */ nullptr, + /* aVendorModel */ nullptr, + /* aVendorSwVersion */ nullptr, + /* aVendorData */ nullptr, + /* aCallback */ nullptr, + /* aContext */ nullptr)); + + joinerIids[1].SetFromExtAddress(joiners[1]->Get().GetId()); + + nexus.AdvanceTime(8 * Time::kOneSecondInMsec); + + for (uint8_t i = 0; i < kNumEnrollers; i++) + { + Ip6::InterfaceIdentifier readIid; + uint16_t joinerRouterRloc; + + message = AsCoapMessagePtr(recvContext[i].mRelayRxMsgs.GetHead()); + + if ((modes[i] & MeshCoP::EnrollerModeTlv::kForwardJoinerRelayRx) == 0) + { + VerifyOrQuit(message == nullptr); + continue; + } + + VerifyOrQuit(message != nullptr); + + VerifyOrQuit(message->ReadType() == Coap::kTypeNonConfirmable); + VerifyOrQuit(message->ReadCode() == Coap::kCodePost); + SuccessOrQuit(Tlv::Find(*message, readIid)); + SuccessOrQuit(Tlv::Find(*message, joinerRouterRloc)); + + VerifyOrQuit(readIid == joinerIids[1]); + VerifyOrQuit(joinerRouterRloc == admitter.Get().GetRloc16()); + } + + joiners[1]->Get().Stop(); + + for (uint8_t i = 0; i < kNumEnrollers; i++) + { + recvContext[i].Clear(); + } + + // - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + Log("Send `EnrollerKeepAlive` message from `enrollers[0]` to revert its mode back to allowing both Joiner and UDP"); + + modes[0] = MeshCoP::EnrollerModeTlv::kForwardJoinerRelayRx | MeshCoP::EnrollerModeTlv::kForwardUdpProxyRx; + + message = + enrollers[0]->Get().AllocateAndInitPriorityConfirmablePostMessage(kUriEnrollerKeepAlive); + VerifyOrQuit(message != nullptr); + + SuccessOrQuit(Tlv::Append(*message, MeshCoP::StateTlv::kAccept)); + SuccessOrQuit(Tlv::Append(*message, modes[0])); + + responseContexts[0].Clear(); + SuccessOrQuit(enrollers[0]->Get().SendMessage(*message, HandleResponse, &responseContexts[0])); + + nexus.AdvanceTime(Time::kOneSecondInMsec); + + VerifyOrQuit(responseContexts[0].mReceived); + VerifyOrQuit(responseContexts[0].mResponseState == MeshCoP::StateTlv::kAccept); + VerifyOrQuit(responseContexts[0].mHasAdmitterState); + + for (uint8_t i = 0; i < kNumEnrollers; i++) + { + recvContext[i].Clear(); + } + + // - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - + Log("Validate that the `enrollers[0]` mode got changed on `admitter`"); + + foundEnrollers.Clear(); + iter.Init(admitter.GetInstance()); + + while (iter.GetNextEnrollerInfo(enrollerInfo) == kErrorNone) + { + uint8_t matchedIndex; + + LogEnroller(enrollerInfo); + + matchedIndex = FindMatchingEnroller(enrollerInfo, kEnrollerIds, foundEnrollers); + + VerifyOrQuit(AsCoreType(&enrollerInfo.mSteeringData) == steeringData); + VerifyOrQuit(enrollerInfo.mMode == modes[matchedIndex]); + + if (matchedIndex == 0) + { + SuccessOrQuit(iter.GetNextJoinerInfo(joinerInfo)); + VerifyOrQuit(AsCoreType(&joinerInfo.mIid) == joinerIids[0]); + LogJoiner(joinerInfo); + } + + VerifyOrQuit(iter.GetNextJoinerInfo(joinerInfo) == kErrorNotFound); + } + + VerifyOrQuit(DidFindAllEnrollers(foundEnrollers)); + + nexus.AdvanceTime(Time::kOneSecondInMsec); + + for (uint8_t i = 0; i < kNumEnrollers; i++) + { + recvContext[i].Clear(); + } + // - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - Log("Send `EnrollerKeepAlive` message from all `enrollers` to maintain the connection");