diff --git a/src/core/common/tlvs.cpp b/src/core/common/tlvs.cpp index 90aef0a47..8dd90695d 100644 --- a/src/core/common/tlvs.cpp +++ b/src/core/common/tlvs.cpp @@ -37,6 +37,9 @@ #include "common/debug.hpp" #include "common/message.hpp" +using ot::Encoding::BigEndian::HostSwap16; +using ot::Encoding::BigEndian::HostSwap32; + namespace ot { uint32_t Tlv::GetSize(void) const @@ -174,14 +177,55 @@ exit: return error; } +otError Tlv::ReadUint8Tlv(const Message &aMessage, uint16_t aOffset, uint8_t &aValue) +{ + return ReadTlv(aMessage, aOffset, &aValue, sizeof(uint8_t)); +} + +otError Tlv::ReadUint16Tlv(const Message &aMessage, uint16_t aOffset, uint16_t &aValue) +{ + otError error; + + SuccessOrExit(error = ReadTlv(aMessage, aOffset, &aValue, sizeof(uint16_t))); + aValue = HostSwap16(aValue); + +exit: + return error; +} + +otError Tlv::ReadUint32Tlv(const Message &aMessage, uint16_t aOffset, uint32_t &aValue) +{ + otError error; + + SuccessOrExit(error = ReadTlv(aMessage, aOffset, &aValue, sizeof(uint32_t))); + aValue = HostSwap32(aValue); + +exit: + return error; +} + +otError Tlv::ReadTlv(const Message &aMessage, uint16_t aOffset, void *aValue, uint8_t aLength) +{ + otError error = OT_ERROR_NONE; + Tlv tlv; + + VerifyOrExit(aMessage.Read(aOffset, sizeof(Tlv), &tlv) == sizeof(Tlv), error = OT_ERROR_PARSE); + VerifyOrExit(!tlv.IsExtended() && (tlv.GetLength() >= aLength), error = OT_ERROR_PARSE); + VerifyOrExit(tlv.GetSize() + aOffset <= aMessage.GetLength(), error = OT_ERROR_PARSE); + + aMessage.Read(aOffset + sizeof(Tlv), aLength, aValue); + +exit: + return error; +} + otError Tlv::FindUint8Tlv(const Message &aMessage, uint8_t aType, uint8_t &aValue) { otError error = OT_ERROR_NONE; - TlvUint8 tlv8; + uint16_t offset; - SuccessOrExit(error = FindTlv(aMessage, aType, sizeof(tlv8), tlv8)); - VerifyOrExit(tlv8.IsValid(), error = OT_ERROR_PARSE); - aValue = tlv8.GetUint8Value(); + SuccessOrExit(error = FindTlvOffset(aMessage, aType, offset)); + error = ReadUint8Tlv(aMessage, offset, aValue); exit: return error; @@ -189,12 +233,11 @@ exit: otError Tlv::FindUint16Tlv(const Message &aMessage, uint8_t aType, uint16_t &aValue) { - otError error = OT_ERROR_NONE; - TlvUint16 tlv16; + otError error = OT_ERROR_NONE; + uint16_t offset; - SuccessOrExit(error = FindTlv(aMessage, aType, sizeof(tlv16), tlv16)); - VerifyOrExit(tlv16.IsValid(), error = OT_ERROR_PARSE); - aValue = tlv16.GetUint16Value(); + SuccessOrExit(error = FindTlvOffset(aMessage, aType, offset)); + error = ReadUint16Tlv(aMessage, offset, aValue); exit: return error; @@ -202,12 +245,11 @@ exit: otError Tlv::FindUint32Tlv(const Message &aMessage, uint8_t aType, uint32_t &aValue) { - otError error = OT_ERROR_NONE; - TlvUint32 tlv32; + otError error = OT_ERROR_NONE; + uint16_t offset; - SuccessOrExit(error = FindTlv(aMessage, aType, sizeof(tlv32), tlv32)); - VerifyOrExit(tlv32.IsValid(), error = OT_ERROR_PARSE); - aValue = tlv32.GetUint32Value(); + SuccessOrExit(error = FindTlvOffset(aMessage, aType, offset)); + error = ReadUint32Tlv(aMessage, offset, aValue); exit: return error; @@ -229,32 +271,23 @@ exit: otError Tlv::AppendUint8Tlv(Message &aMessage, uint8_t aType, uint8_t aValue) { - TlvUint8 tlv8; + uint8_t value8 = aValue; - tlv8.Init(aType); - tlv8.SetUint8Value(aValue); - - return tlv8.AppendTo(aMessage); + return AppendTlv(aMessage, aType, &value8, sizeof(uint8_t)); } otError Tlv::AppendUint16Tlv(Message &aMessage, uint8_t aType, uint16_t aValue) { - TlvUint16 tlv16; + uint16_t value16 = HostSwap16(aValue); - tlv16.Init(aType); - tlv16.SetUint16Value(aValue); - - return tlv16.AppendTo(aMessage); + return AppendTlv(aMessage, aType, &value16, sizeof(uint16_t)); } otError Tlv::AppendUint32Tlv(Message &aMessage, uint8_t aType, uint32_t aValue) { - TlvUint32 tlv32; + uint32_t value32 = HostSwap32(aValue); - tlv32.Init(aType); - tlv32.SetUint32Value(aValue); - - return tlv32.AppendTo(aMessage); + return AppendTlv(aMessage, aType, &value32, sizeof(uint32_t)); } otError Tlv::AppendTlv(Message &aMessage, uint8_t aType, const void *aValue, uint8_t aLength) diff --git a/src/core/common/tlvs.hpp b/src/core/common/tlvs.hpp index cebf55e7b..d3039c88f 100644 --- a/src/core/common/tlvs.hpp +++ b/src/core/common/tlvs.hpp @@ -45,7 +45,6 @@ namespace ot { using ot::Encoding::BigEndian::HostSwap16; -using ot::Encoding::BigEndian::HostSwap32; class Message; @@ -177,6 +176,59 @@ public: */ otError AppendTo(Message &aMessage) const; + /** + * This static method reads a TLV from a message at a given offset with TLV's value as an `uint8_t`. + * + * @param[in] aMessage The message to read from. + * @param[in] aOffset The offset into the message pointing to the start of the TLV. + * @param[out] aValue A reference to a `uint8_t` to output the TLV's value. + * + * @retval OT_ERROR_NONE Successfully read the TLV and updated @p aValue. + * @retval OT_ERROR_PARSE The TLV was not well-formed and could not be parsed. + * + */ + static otError ReadUint8Tlv(const Message &aMessage, uint16_t aOffset, uint8_t &aValue); + + /** + * This static method reads a TLV from a message at a given offset with TLV's value as an `uint16_t`. + * + * @param[in] aMessage The message to read from. + * @param[in] aOffset The offset into the message pointing to the start of the TLV. + * @param[out] aValue A reference to a `uint16_t` to output the TLV's value. + * + * @retval OT_ERROR_NONE Successfully read the TLV and updated @p aValue. + * @retval OT_ERROR_PARSE The TLV was not well-formed and could not be parsed. + * + */ + static otError ReadUint16Tlv(const Message &aMessage, uint16_t aOffset, uint16_t &aValue); + + /** + * This static method reads a TLV from a message at a given offset with TLV's value as an `uint32_t`. + * + * @param[in] aMessage The message to read from. + * @param[in] aOffset The offset into the message pointing to the start of the TLV. + * @param[out] aValue A reference to a `uint32_t` to output the TLV's value. + * + * @retval OT_ERROR_NONE Successfully read the TLV and updated @p aValue. + * @retval OT_ERROR_PARSE The TLV was not well-formed and could not be parsed. + * + */ + static otError ReadUint32Tlv(const Message &aMessage, uint16_t aOffset, uint32_t &aValue); + + /** + * This static method reads a TLV in a message at a given offset expecting a minimum length for the value. + * + * @param[in] aMessage The message to read from. + * @param[in] aOffset The offset into the message pointing to the start of the TLV. + * @param[out] aValue A buffer to output the TLV's value, must contain (at least) @p aMinLength bytes. + * @param[in] aMinLength The minimum expected length of TLV and number of bytes to copy into @p aValue buffer. + * + * @retval OT_ERROR_NONE Successfully read the TLV and copied @p aMinLength into @p aValue. + * @retval OT_ERROR_PARSE The TLV was not well-formed and could not be parsed. + * + */ + static otError ReadTlv(const Message &aMessage, uint16_t aOffset, void *aValue, uint8_t aMinLength); + /** * This static method reads the requested TLV out of @p aMessage. * @@ -409,153 +461,6 @@ private: uint16_t mLength; } OT_TOOL_PACKED_END; -/** - * This class implements a simple TLV with a `uint8_t` value. - * - */ -OT_TOOL_PACKED_BEGIN -class TlvUint8 : public Tlv -{ -public: - /** - * This method initializes the TLV. - * - * @param[in] aType The Type value. - * - */ - void Init(uint8_t aType) - { - SetType(aType); - SetLength(sizeof(*this) - sizeof(Tlv)); - } - - /** - * This method indicates whether or not the TLV appears to be well-formed. - * - * @retval TRUE If the TLV appears to be well-formed. - * @retval FALSE If the TLV does not appear to be well-formed. - * - */ - bool IsValid(void) const { return GetLength() >= sizeof(*this) - sizeof(Tlv); } - - /** - * This method returns the `uint8_t` value. - * - * @returns The `uint8_t` value. - * - */ - uint8_t GetUint8Value(void) const { return mValue; } - - /** - * This method sets the `uint8_t` value. - * - * @param[in] aValue The `uint8_t` value. - * - */ - void SetUint8Value(uint8_t aValue) { mValue = aValue; } - -private: - uint8_t mValue; -} OT_TOOL_PACKED_END; - -/** - * This class implements a simple TLV with a `uint16_t` value. - * - */ -OT_TOOL_PACKED_BEGIN -class TlvUint16 : public Tlv -{ -public: - /** - * This method initializes the TLV. - * - * @param[in] aType The Type value. - * - */ - void Init(uint8_t aType) - { - SetType(aType); - SetLength(sizeof(*this) - sizeof(Tlv)); - } - - /** - * This method indicates whether or not the TLV appears to be well-formed. - * - * @retval TRUE If the TLV appears to be well-formed. - * @retval FALSE If the TLV does not appear to be well-formed. - * - */ - bool IsValid(void) const { return GetLength() >= sizeof(*this) - sizeof(Tlv); } - - /** - * This method returns the `uint16_t` value. - * - * @returns The `uint16_t` value. - * - */ - uint16_t GetUint16Value(void) const { return HostSwap16(mValue); } - - /** - * This method sets the `uint16_t` value. - * - * @param[in] aValue The `uint16_t` value. - * - */ - void SetUint16Value(uint16_t aValue) { mValue = HostSwap16(aValue); } - -private: - uint16_t mValue; -} OT_TOOL_PACKED_END; - -/** - * This class implements a simple TLV with a `uint32_t` value. - * - */ -OT_TOOL_PACKED_BEGIN -class TlvUint32 : public Tlv -{ -public: - /** - * This method initializes the TLV. - * - * @param[in] aType The Type value. - * - */ - void Init(uint8_t aType) - { - SetType(aType); - SetLength(sizeof(*this) - sizeof(Tlv)); - } - - /** - * This method indicates whether or not the TLV appears to be well-formed. - * - * @retval TRUE If the TLV appears to be well-formed. - * @retval FALSE If the TLV does not appear to be well-formed. - * - */ - bool IsValid(void) const { return GetLength() >= sizeof(*this) - sizeof(Tlv); } - - /** - * This method returns the `uint32_t` value. - * - * @returns The `uint32_t` value. - * - */ - uint32_t GetUint32Value(void) const { return HostSwap32(mValue); } - - /** - * This method sets the `uint32_t` value. - * - * @param[in] aValue The `uint32_t` value. - * - */ - void SetUint32Value(uint32_t aValue) { mValue = HostSwap32(aValue); } - -private: - uint32_t mValue; -} OT_TOOL_PACKED_END; - } // namespace ot #endif // TLVS_HPP_ diff --git a/src/core/thread/mle.cpp b/src/core/thread/mle.cpp index 103ecab17..6ed9b1d7a 100644 --- a/src/core/thread/mle.cpp +++ b/src/core/thread/mle.cpp @@ -3857,10 +3857,8 @@ void Mle::HandleDiscoveryResponse(const Message &aMessage, const Ip6::MessageInf Tlv tlv; MeshCoP::Tlv meshcopTlv; MeshCoP::DiscoveryResponseTlv discoveryResponse; - MeshCoP::ExtendedPanIdTlv extPanId; MeshCoP::NetworkNameTlv networkName; MeshCoP::SteeringDataTlv steeringData; - MeshCoP::JoinerUdpPortTlv JoinerUdpPort; otActiveScanResult result; uint16_t offset; uint16_t end; @@ -3899,9 +3897,7 @@ void Mle::HandleDiscoveryResponse(const Message &aMessage, const Ip6::MessageInf break; case MeshCoP::Tlv::kExtendedPanId: - aMessage.Read(offset, sizeof(extPanId), &extPanId); - VerifyOrExit(extPanId.IsValid(), error = OT_ERROR_PARSE); - result.mExtendedPanId = extPanId.GetExtendedPanId(); + SuccessOrExit(error = Tlv::ReadTlv(aMessage, offset, &result.mExtendedPanId, sizeof(Mac::ExtendedPanId))); break; case MeshCoP::Tlv::kNetworkName: @@ -3927,9 +3923,7 @@ void Mle::HandleDiscoveryResponse(const Message &aMessage, const Ip6::MessageInf break; case MeshCoP::Tlv::kJoinerUdpPort: - aMessage.Read(offset, sizeof(JoinerUdpPort), &JoinerUdpPort); - VerifyOrExit(JoinerUdpPort.IsValid(), error = OT_ERROR_PARSE); - result.mJoinerUdpPort = JoinerUdpPort.GetUdpPort(); + SuccessOrExit(error = Tlv::ReadUint16Tlv(aMessage, offset, result.mJoinerUdpPort)); break; default: diff --git a/src/core/thread/mle_router.cpp b/src/core/thread/mle_router.cpp index 58b2dc2dd..fd35b4dfe 100644 --- a/src/core/thread/mle_router.cpp +++ b/src/core/thread/mle_router.cpp @@ -2785,7 +2785,7 @@ void MleRouter::HandleDiscoveryRequest(const Message &aMessage, const Ip6::Messa Tlv tlv; MeshCoP::Tlv meshcopTlv; MeshCoP::DiscoveryRequestTlv discoveryRequest; - MeshCoP::ExtendedPanIdTlv extPanId; + Mac::ExtendedPanId extPanId; uint16_t offset; uint16_t end; @@ -2829,9 +2829,8 @@ void MleRouter::HandleDiscoveryRequest(const Message &aMessage, const Ip6::Messa break; case MeshCoP::Tlv::kExtendedPanId: - aMessage.Read(offset, sizeof(extPanId), &extPanId); - VerifyOrExit(extPanId.IsValid(), error = OT_ERROR_PARSE); - VerifyOrExit(Get().GetExtendedPanId() != extPanId.GetExtendedPanId(), error = OT_ERROR_DROP); + SuccessOrExit(error = Tlv::ReadTlv(aMessage, offset, &extPanId, sizeof(extPanId))); + VerifyOrExit(Get().GetExtendedPanId() != extPanId, error = OT_ERROR_DROP); break; diff --git a/src/core/thread/network_diagnostic.cpp b/src/core/thread/network_diagnostic.cpp index 0d2f0aaee..650a600b5 100644 --- a/src/core/thread/network_diagnostic.cpp +++ b/src/core/thread/network_diagnostic.cpp @@ -770,52 +770,26 @@ otError NetworkDiagnostic::GetNextDiagTlv(const Coap::Message &aMessage, switch (tlv.GetType()) { case NetworkDiagnosticTlv::kExtMacAddress: - { - ExtMacAddressTlv extMacAddr; - - tlvTotalLength = sizeof(extMacAddr); - VerifyOrExit(aMessage.Read(offset, tlvTotalLength, &extMacAddr) == tlvTotalLength, error = OT_ERROR_PARSE); - VerifyOrExit(extMacAddr.IsValid(), error = OT_ERROR_PARSE); - - aNetworkDiagTlv.mData.mExtAddress = *extMacAddr.GetMacAddr(); + SuccessOrExit( + error = Tlv::ReadTlv(aMessage, offset, &aNetworkDiagTlv.mData.mExtAddress, sizeof(Mac::ExtAddress))); break; - } case NetworkDiagnosticTlv::kAddress16: - { - Address16Tlv addr16; - - tlvTotalLength = sizeof(addr16); - VerifyOrExit(aMessage.Read(offset, tlvTotalLength, &addr16) == tlvTotalLength, error = OT_ERROR_PARSE); - VerifyOrExit(addr16.IsValid(), error = OT_ERROR_PARSE); - - aNetworkDiagTlv.mData.mAddr16 = addr16.GetRloc16(); + SuccessOrExit(error = Tlv::ReadUint16Tlv(aMessage, offset, aNetworkDiagTlv.mData.mAddr16)); break; - } case NetworkDiagnosticTlv::kMode: { - ModeTlv linkMode; + uint8_t mode; - tlvTotalLength = sizeof(linkMode); - VerifyOrExit(aMessage.Read(offset, tlvTotalLength, &linkMode) == tlvTotalLength, error = OT_ERROR_PARSE); - VerifyOrExit(linkMode.IsValid(), error = OT_ERROR_PARSE); - - ParseMode(linkMode.GetMode(), aNetworkDiagTlv.mData.mMode); + SuccessOrExit(error = Tlv::ReadUint8Tlv(aMessage, offset, mode)); + ParseMode(Mle::DeviceMode(mode), aNetworkDiagTlv.mData.mMode); break; } case NetworkDiagnosticTlv::kTimeout: - { - TimeoutTlv timeout; - - tlvTotalLength = sizeof(timeout); - VerifyOrExit(aMessage.Read(offset, tlvTotalLength, &timeout) == tlvTotalLength, error = OT_ERROR_PARSE); - VerifyOrExit(timeout.IsValid(), error = OT_ERROR_PARSE); - - aNetworkDiagTlv.mData.mTimeout = timeout.GetTimeout(); + SuccessOrExit(error = Tlv::ReadUint32Tlv(aMessage, offset, aNetworkDiagTlv.mData.mTimeout)); break; - } case NetworkDiagnosticTlv::kConnectivity: { @@ -899,30 +873,12 @@ otError NetworkDiagnostic::GetNextDiagTlv(const Coap::Message &aMessage, } case NetworkDiagnosticTlv::kBatteryLevel: - { - BatteryLevelTlv batteryLevel; - - tlvTotalLength = sizeof(BatteryLevelTlv); - VerifyOrExit(aMessage.Read(offset, tlvTotalLength, &batteryLevel) == tlvTotalLength, - error = OT_ERROR_PARSE); - VerifyOrExit(batteryLevel.IsValid(), error = OT_ERROR_PARSE); - - aNetworkDiagTlv.mData.mBatteryLevel = batteryLevel.GetBatteryLevel(); + SuccessOrExit(error = Tlv::ReadUint8Tlv(aMessage, offset, aNetworkDiagTlv.mData.mBatteryLevel)); break; - } case NetworkDiagnosticTlv::kSupplyVoltage: - { - SupplyVoltageTlv supplyVoltage; - - tlvTotalLength = sizeof(SupplyVoltageTlv); - VerifyOrExit(aMessage.Read(offset, tlvTotalLength, &supplyVoltage) == tlvTotalLength, - error = OT_ERROR_PARSE); - VerifyOrExit(supplyVoltage.IsValid(), error = OT_ERROR_PARSE); - - aNetworkDiagTlv.mData.mSupplyVoltage = supplyVoltage.GetSupplyVoltage(); + SuccessOrExit(error = Tlv::ReadUint16Tlv(aMessage, offset, aNetworkDiagTlv.mData.mSupplyVoltage)); break; - } case NetworkDiagnosticTlv::kChildTable: { @@ -955,17 +911,8 @@ otError NetworkDiagnostic::GetNextDiagTlv(const Coap::Message &aMessage, } case NetworkDiagnosticTlv::kMaxChildTimeout: - { - MaxChildTimeoutTlv maxChildTimeout; - - tlvTotalLength = sizeof(maxChildTimeout); - VerifyOrExit(aMessage.Read(offset, tlvTotalLength, &maxChildTimeout) == tlvTotalLength, - error = OT_ERROR_PARSE); - VerifyOrExit(maxChildTimeout.IsValid(), error = OT_ERROR_PARSE); - - aNetworkDiagTlv.mData.mMaxChildTimeout = maxChildTimeout.GetTimeout(); + SuccessOrExit(error = Tlv::ReadUint32Tlv(aMessage, offset, aNetworkDiagTlv.mData.mMaxChildTimeout)); break; - } default: // Ignore unrecognized Network Diagnostic TLV silently and