From 977ac2174390793608d594d8715eab1acf97e6ed Mon Sep 17 00:00:00 2001 From: Abtin Keshavarzian Date: Mon, 5 Jun 2023 12:31:26 -0700 Subject: [PATCH] [sntp-client] smaller enhancements (#9125) This commit contains the following changes in the `Sntp::Client` class: - The `Header` and `QueryMetadata` classes are moved to be internal (nested) classes of `Client`. - The `Header` constructor is removed and replaced with an `Init()` method. This is to avoid unnecessary initialization of a `Header` instances when reading it from a received message. - The `QueryMetadata` constructor has been removed. - The `QueryMetadata` uses `Callback<>` for `ResponseHandler`. --- src/core/net/sntp_client.cpp | 54 +--- src/core/net/sntp_client.hpp | 515 ++++++++--------------------------- 2 files changed, 122 insertions(+), 447 deletions(-) diff --git a/src/core/net/sntp_client.cpp b/src/core/net/sntp_client.cpp index e93effa7c..bc79e3d9b 100644 --- a/src/core/net/sntp_client.cpp +++ b/src/core/net/sntp_client.cpp @@ -50,49 +50,6 @@ namespace Sntp { RegisterLogModule("SntpClnt"); -Header::Header(void) - : mFlags(kNtpVersion << kVersionOffset | kModeClient << kModeOffset) - , mStratum(0) - , mPoll(0) - , mPrecision(0) - , mRootDelay(0) - , mRootDispersion(0) - , mReferenceId(0) - , mReferenceTimestampSeconds(0) - , mReferenceTimestampFraction(0) - , mOriginateTimestampSeconds(0) - , mOriginateTimestampFraction(0) - , mReceiveTimestampSeconds(0) - , mReceiveTimestampFraction(0) - , mTransmitTimestampSeconds(0) - , mTransmitTimestampFraction(0) -{ -} - -QueryMetadata::QueryMetadata(void) - : mTransmitTimestamp(0) - , mResponseHandler(nullptr) - , mResponseContext(nullptr) - , mTransmissionTime(0) - , mDestinationPort(0) - , mRetransmissionCount(0) -{ - mSourceAddress.Clear(); - mDestinationAddress.Clear(); -} - -QueryMetadata::QueryMetadata(otSntpResponseHandler aHandler, void *aContext) - : mTransmitTimestamp(0) - , mResponseHandler(aHandler) - , mResponseContext(aContext) - , mTransmissionTime(0) - , mDestinationPort(0) - , mRetransmissionCount(0) -{ - mSourceAddress.Clear(); - mDestinationAddress.Clear(); -} - Client::Client(Instance &aInstance) : mSocket(aInstance) , mRetransmissionTimer(aInstance) @@ -127,7 +84,7 @@ Error Client::Stop(void) Error Client::Query(const otSntpQuery *aQuery, otSntpResponseHandler aHandler, void *aContext) { Error error; - QueryMetadata queryMetadata(aHandler, aContext); + QueryMetadata queryMetadata; Message *message = nullptr; Message *messageCopy = nullptr; Header header; @@ -135,6 +92,8 @@ Error Client::Query(const otSntpQuery *aQuery, otSntpResponseHandler aHandler, v VerifyOrExit(aQuery->mMessageInfo != nullptr, error = kErrorInvalidArgs); + header.Init(); + // Originate timestamp is used only as a unique token. header.SetTransmitTimestampSeconds(TimerMilli::GetNow().GetValue() / 1000 + kTimeAt1970); @@ -142,6 +101,7 @@ Error Client::Query(const otSntpQuery *aQuery, otSntpResponseHandler aHandler, v messageInfo = AsCoreTypePtr(aQuery->mMessageInfo); + queryMetadata.mResponseHandler.Set(aHandler, aContext); queryMetadata.mTransmitTimestamp = header.GetTransmitTimestampSeconds(); queryMetadata.mTransmissionTime = TimerMilli::GetNow() + kResponseTimeout; queryMetadata.mSourceAddress = messageInfo->GetSockAddr(); @@ -262,11 +222,7 @@ void Client::FinalizeSntpTransaction(Message &aQuery, Error aResult) { DequeueMessage(aQuery); - - if (aQueryMetadata.mResponseHandler != nullptr) - { - aQueryMetadata.mResponseHandler(aQueryMetadata.mResponseContext, aTime, aResult); - } + aQueryMetadata.mResponseHandler.InvokeIfSet(aTime, aResult); } void Client::HandleRetransmissionTimer(void) diff --git a/src/core/net/sntp_client.hpp b/src/core/net/sntp_client.hpp index 6479ac341..3a513aa4e 100644 --- a/src/core/net/sntp_client.hpp +++ b/src/core/net/sntp_client.hpp @@ -35,6 +35,7 @@ #include +#include "common/clearable.hpp" #include "common/message.hpp" #include "common/non_copyable.hpp" #include "common/timer.hpp" @@ -51,403 +52,6 @@ namespace Sntp { using ot::Encoding::BigEndian::HostSwap32; -/** - * Implements SNTP header generation and parsing. - * - */ -OT_TOOL_PACKED_BEGIN -class Header -{ -public: - /** - * Default constructor for SNTP Header. - * - */ - Header(void); - - /** - * Defines supported SNTP modes. - * - */ - enum Mode : uint8_t - { - kModeClient = 3, - kModeServer = 4, - }; - - static constexpr uint8_t kKissCodeLength = 4; ///< Length of the kiss code in ASCII format - - /** - * Returns the flags field value. - * - * @returns Value of the flags field (LI, VN and Mode). - * - */ - uint8_t GetFlags(void) const { return mFlags; } - - /** - * Sets the flags field. - * - * @param[in] aFlags The value of the flags field. - * - */ - void SetFlags(uint8_t aFlags) { mFlags = aFlags; } - - /** - * Returns the SNTP operational mode. - * - * @returns SNTP operational mode. - * - */ - Mode GetMode(void) const { return static_cast((mFlags & kModeMask) >> kModeOffset); } - - /** - * Returns the packet stratum field value. - * - * @returns Value of the packet stratum. - * - */ - uint8_t GetStratum(void) const { return mStratum; } - - /** - * Sets the packet stratum field value. - * - * @param[in] aStratum The value of the packet stratum field. - * - */ - void SetStratum(uint8_t aStratum) { mStratum = aStratum; } - - /** - * Returns the poll field value. - * - * @returns Value of the poll field. - * - */ - uint8_t GetPoll(void) const { return mPoll; } - - /** - * Sets the poll field. - * - * @param[in] aPoll The value of the poll field. - * - */ - void SetPoll(uint8_t aPoll) { mPoll = aPoll; } - - /** - * Returns the precision field value. - * - * @returns Value of the precision field. - * - */ - uint8_t GetPrecision(void) const { return mPrecision; } - - /** - * Sets the precision field. - * - * @param[in] aPrecision The value of the precision field. - * - */ - void SetPrecision(uint8_t aPrecision) { mPrecision = aPrecision; } - - /** - * Returns the root delay field value. - * - * @returns Value of the root delay field. - * - */ - uint32_t GetRootDelay(void) const { return HostSwap32(mRootDelay); } - - /** - * Sets the root delay field. - * - * @param[in] aRootDelay The value of the root delay field. - * - */ - void SetRootDelay(uint32_t aRootDelay) { mRootDelay = HostSwap32(aRootDelay); } - - /** - * Returns the root dispersion field value. - * - * @returns Value of the root dispersion field. - * - */ - uint32_t GetRootDispersion(void) const { return HostSwap32(mRootDispersion); } - - /** - * Sets the root dispersion field. - * - * @param[in] aRootDispersion The value of the root dispersion field. - * - */ - void SetRootDispersion(uint32_t aRootDispersion) { mRootDispersion = HostSwap32(aRootDispersion); } - - /** - * Returns the reference identifier field value. - * - * @returns Value of the reference identifier field. - * - */ - uint32_t GetReferenceId(void) const { return HostSwap32(mReferenceId); } - - /** - * Sets the reference identifier field. - * - * @param[in] aReferenceId The value of the reference identifier field. - * - */ - void SetReferenceId(uint32_t aReferenceId) { mReferenceId = HostSwap32(aReferenceId); } - - /** - * Returns the kiss code in ASCII format. - * - * @returns Value of the reference identifier field in ASCII format. - * - */ - char *GetKissCode(void) { return reinterpret_cast(&mReferenceId); } - - /** - * Returns the reference timestamp seconds field. - * - * @returns Value of the reference timestamp seconds field. - * - */ - uint32_t GetReferenceTimestampSeconds(void) const { return HostSwap32(mReferenceTimestampSeconds); } - - /** - * Sets the reference timestamp seconds field. - * - * @param[in] aReferenceTimestampSeconds Value of the reference timestamp seconds field. - * - */ - void SetReferenceTimestampSeconds(uint32_t aReferenceTimestampSeconds) - { - mReferenceTimestampSeconds = HostSwap32(aReferenceTimestampSeconds); - } - - /** - * Returns the reference timestamp fraction field. - * - * @returns Value of the reference timestamp fraction field. - * - */ - uint32_t GetReferenceTimestampFraction(void) const { return HostSwap32(mReferenceTimestampFraction); } - - /** - * Sets the reference timestamp fraction field. - * - * @param[in] aReferenceTimestampFraction Value of the reference timestamp fraction field. - * - */ - void SetReferenceTimestampFraction(uint32_t aReferenceTimestampFraction) - { - mReferenceTimestampFraction = HostSwap32(aReferenceTimestampFraction); - } - - /** - * Returns the originate timestamp seconds field. - * - * @returns Value of the originate timestamp seconds field. - * - */ - uint32_t GetOriginateTimestampSeconds(void) const { return HostSwap32(mOriginateTimestampSeconds); } - - /** - * Sets the originate timestamp seconds field. - * - * @param[in] aOriginateTimestampSeconds Value of the originate timestamp seconds field. - * - */ - void SetOriginateTimestampSeconds(uint32_t aOriginateTimestampSeconds) - { - mOriginateTimestampSeconds = HostSwap32(aOriginateTimestampSeconds); - } - - /** - * Returns the originate timestamp fraction field. - * - * @returns Value of the originate timestamp fraction field. - * - */ - uint32_t GetOriginateTimestampFraction(void) const { return HostSwap32(mOriginateTimestampFraction); } - - /** - * Sets the originate timestamp fraction field. - * - * @param[in] aOriginateTimestampFraction Value of the originate timestamp fraction field. - * - */ - void SetOriginateTimestampFraction(uint32_t aOriginateTimestampFraction) - { - mOriginateTimestampFraction = HostSwap32(aOriginateTimestampFraction); - } - - /** - * Returns the receive timestamp seconds field. - * - * @returns Value of the receive timestamp seconds field. - * - */ - uint32_t GetReceiveTimestampSeconds(void) const { return HostSwap32(mReceiveTimestampSeconds); } - - /** - * Sets the receive timestamp seconds field. - * - * @param[in] aReceiveTimestampSeconds Value of the receive timestamp seconds field. - * - */ - void SetReceiveTimestampSeconds(uint32_t aReceiveTimestampSeconds) - { - mReceiveTimestampSeconds = HostSwap32(aReceiveTimestampSeconds); - } - - /** - * Returns the receive timestamp fraction field. - * - * @returns Value of the receive timestamp fraction field. - * - */ - uint32_t GetReceiveTimestampFraction(void) const { return HostSwap32(mReceiveTimestampFraction); } - - /** - * Sets the receive timestamp fraction field. - * - * @param[in] aReceiveTimestampFraction Value of the receive timestamp fraction field. - * - */ - void SetReceiveTimestampFraction(uint32_t aReceiveTimestampFraction) - { - mReceiveTimestampFraction = HostSwap32(aReceiveTimestampFraction); - } - - /** - * Returns the transmit timestamp seconds field. - * - * @returns Value of the transmit timestamp seconds field. - * - */ - uint32_t GetTransmitTimestampSeconds(void) const { return HostSwap32(mTransmitTimestampSeconds); } - - /** - * Sets the transmit timestamp seconds field. - * - * @param[in] aTransmitTimestampSeconds Value of the transmit timestamp seconds field. - * - */ - void SetTransmitTimestampSeconds(uint32_t aTransmitTimestampSeconds) - { - mTransmitTimestampSeconds = HostSwap32(aTransmitTimestampSeconds); - } - - /** - * Returns the transmit timestamp fraction field. - * - * @returns Value of the transmit timestamp fraction field. - * - */ - uint32_t GetTransmitTimestampFraction(void) const { return HostSwap32(mTransmitTimestampFraction); } - - /** - * Sets the transmit timestamp fraction field. - * - * @param[in] aTransmitTimestampFraction Value of the transmit timestamp fraction field. - * - */ - void SetTransmitTimestampFraction(uint32_t aTransmitTimestampFraction) - { - mTransmitTimestampFraction = HostSwap32(aTransmitTimestampFraction); - } - -private: - static constexpr uint8_t kNtpVersion = 4; // Current NTP version. - static constexpr uint8_t kLeapOffset = 6; // Leap Indicator field offset. - static constexpr uint8_t kLeapMask = 0x03 << kLeapOffset; // Leap Indicator field mask. - static constexpr uint8_t kVersionOffset = 3; // Version field offset. - static constexpr uint8_t kVersionMask = 0x07 << kVersionOffset; // Version field mask. - static constexpr uint8_t kModeOffset = 0; // Mode field offset. - static constexpr uint8_t kModeMask = 0x07 << kModeOffset; // Mode filed mask. - - uint8_t mFlags; // SNTP flags: LI Leap Indicator, VN Version Number and Mode. - uint8_t mStratum; // Packet Stratum. - uint8_t mPoll; // Maximum interval between successive messages, in log2 seconds. - uint8_t mPrecision; // The precision of the system clock, in log2 seconds. - uint32_t mRootDelay; // Total round-trip delay to the reference clock, in NTP short format. - uint32_t mRootDispersion; // Total dispersion to the reference clock. - uint32_t mReferenceId; // ID identifying the particular server or reference clock. - uint32_t mReferenceTimestampSeconds; // Time the system clock was last set or corrected (NTP format). - uint32_t mReferenceTimestampFraction; // Fraction part of above value. - uint32_t mOriginateTimestampSeconds; // Time at the client when the request departed for the server (NTP format). - uint32_t mOriginateTimestampFraction; // Fraction part of above value. - uint32_t mReceiveTimestampSeconds; // Time at the server when the request arrived from the client (NTP format). - uint32_t mReceiveTimestampFraction; // Fraction part of above value. - uint32_t mTransmitTimestampSeconds; // Time at the server when the response left for the client (NTP format). - uint32_t mTransmitTimestampFraction; // Fraction part of above value. -} OT_TOOL_PACKED_END; - -/** - * Implements metadata required for SNTP retransmission. - * - */ -class QueryMetadata -{ - friend class Client; - -public: - /** - * Default constructor for the object. - * - */ - QueryMetadata(void); - - /** - * Initializes the object with specific values. - * - * @param[in] aHandler Pointer to a handler function for the response. - * @param[in] aContext Context for the handler function. - * - */ - QueryMetadata(otSntpResponseHandler aHandler, void *aContext); - - /** - * Appends request data to the message. - * - * @param[in] aMessage A reference to the message. - * - * @retval kErrorNone Successfully appended the bytes. - * @retval kErrorNoBufs Insufficient available buffers to grow the message. - * - */ - Error AppendTo(Message &aMessage) const { return aMessage.Append(*this); } - - /** - * Reads request data from the message. - * - * @param[in] aMessage A reference to the message. - * - */ - void ReadFrom(const Message &aMessage) - { - SuccessOrAssert(aMessage.Read(aMessage.GetLength() - sizeof(*this), *this)); - } - - /** - * Updates request data in the message. - * - * @param[in] aMessage A reference to the message. - * - */ - void UpdateIn(Message &aMessage) const { aMessage.Write(aMessage.GetLength() - sizeof(*this), *this); } - -private: - uint32_t mTransmitTimestamp; ///< Time at the client when the request departed for the server. - otSntpResponseHandler mResponseHandler; ///< A function pointer that is called on response reception. - void *mResponseContext; ///< A pointer to arbitrary context information. - TimeMilli mTransmissionTime; ///< Time when the timer should shoot for this message. - Ip6::Address mSourceAddress; ///< IPv6 address of the message source. - Ip6::Address mDestinationAddress; ///< IPv6 address of the message destination. - uint16_t mDestinationPort; ///< UDP port of the message destination. - uint8_t mRetransmissionCount; ///< Number of retransmissions. -}; - /** * Implements SNTP client. * @@ -455,6 +59,8 @@ private: class Client : private NonCopyable { public: + typedef otSntpResponseHandler ResponseHandler; ///< Response handler callback. + /** * Initializes the object. * @@ -507,7 +113,7 @@ public: * @retval kErrorInvalidArgs Invalid arguments supplied. * */ - Error Query(const otSntpQuery *aQuery, otSntpResponseHandler aHandler, void *aContext); + Error Query(const otSntpQuery *aQuery, ResponseHandler aHandler, void *aContext); private: static constexpr uint32_t kTimeAt1970 = 2208988800UL; // num seconds between 1st Jan 1900 and 1st Jan 1970. @@ -515,6 +121,119 @@ private: static constexpr uint32_t kResponseTimeout = OPENTHREAD_CONFIG_SNTP_CLIENT_RESPONSE_TIMEOUT; static constexpr uint8_t kMaxRetransmit = OPENTHREAD_CONFIG_SNTP_CLIENT_MAX_RETRANSMIT; + OT_TOOL_PACKED_BEGIN + class Header : public Clearable
+ { + public: + enum Mode : uint8_t + { + kModeClient = 3, + kModeServer = 4, + }; + + static constexpr uint8_t kKissCodeLength = 4; + + void Init(void) + { + Clear(); + mFlags = (kNtpVersion << kVersionOffset | kModeClient << kModeOffset); + } + + uint8_t GetFlags(void) const { return mFlags; } + void SetFlags(uint8_t aFlags) { mFlags = aFlags; } + + Mode GetMode(void) const { return static_cast((mFlags & kModeMask) >> kModeOffset); } + + uint8_t GetStratum(void) const { return mStratum; } + void SetStratum(uint8_t aStratum) { mStratum = aStratum; } + + uint8_t GetPoll(void) const { return mPoll; } + void SetPoll(uint8_t aPoll) { mPoll = aPoll; } + + uint8_t GetPrecision(void) const { return mPrecision; } + void SetPrecision(uint8_t aPrecision) { mPrecision = aPrecision; } + + uint32_t GetRootDelay(void) const { return HostSwap32(mRootDelay); } + void SetRootDelay(uint32_t aRootDelay) { mRootDelay = HostSwap32(aRootDelay); } + + uint32_t GetRootDispersion(void) const { return HostSwap32(mRootDispersion); } + void SetRootDispersion(uint32_t aRootDispersion) { mRootDispersion = HostSwap32(aRootDispersion); } + + uint32_t GetReferenceId(void) const { return HostSwap32(mReferenceId); } + void SetReferenceId(uint32_t aReferenceId) { mReferenceId = HostSwap32(aReferenceId); } + + char *GetKissCode(void) { return reinterpret_cast(&mReferenceId); } + + uint32_t GetReferenceTimestampSeconds(void) const { return HostSwap32(mReferenceTimestampSeconds); } + void SetReferenceTimestampSeconds(uint32_t aTimestamp) { mReferenceTimestampSeconds = HostSwap32(aTimestamp); } + + uint32_t GetReferenceTimestampFraction(void) const { return HostSwap32(mReferenceTimestampFraction); } + void SetReferenceTimestampFraction(uint32_t aFraction) { mReferenceTimestampFraction = HostSwap32(aFraction); } + + uint32_t GetOriginateTimestampSeconds(void) const { return HostSwap32(mOriginateTimestampSeconds); } + void SetOriginateTimestampSeconds(uint32_t aTimestamp) { mOriginateTimestampSeconds = HostSwap32(aTimestamp); } + + uint32_t GetOriginateTimestampFraction(void) const { return HostSwap32(mOriginateTimestampFraction); } + void SetOriginateTimestampFraction(uint32_t aFraction) { mOriginateTimestampFraction = HostSwap32(aFraction); } + + uint32_t GetReceiveTimestampSeconds(void) const { return HostSwap32(mReceiveTimestampSeconds); } + void SetReceiveTimestampSeconds(uint32_t aTimestamp) { mReceiveTimestampSeconds = HostSwap32(aTimestamp); } + + uint32_t GetReceiveTimestampFraction(void) const { return HostSwap32(mReceiveTimestampFraction); } + void SetReceiveTimestampFraction(uint32_t aFraction) { mReceiveTimestampFraction = HostSwap32(aFraction); } + + uint32_t GetTransmitTimestampSeconds(void) const { return HostSwap32(mTransmitTimestampSeconds); } + void SetTransmitTimestampSeconds(uint32_t aTimestamp) { mTransmitTimestampSeconds = HostSwap32(aTimestamp); } + + uint32_t GetTransmitTimestampFraction(void) const { return HostSwap32(mTransmitTimestampFraction); } + void SetTransmitTimestampFraction(uint32_t aFraction) { mTransmitTimestampFraction = HostSwap32(aFraction); } + + private: + static constexpr uint8_t kNtpVersion = 4; // Current NTP version. + static constexpr uint8_t kLeapOffset = 6; // Leap Indicator field offset. + static constexpr uint8_t kLeapMask = 0x03 << kLeapOffset; // Leap Indicator field mask. + static constexpr uint8_t kVersionOffset = 3; // Version field offset. + static constexpr uint8_t kVersionMask = 0x07 << kVersionOffset; // Version field mask. + static constexpr uint8_t kModeOffset = 0; // Mode field offset. + static constexpr uint8_t kModeMask = 0x07 << kModeOffset; // Mode filed mask. + + uint8_t mFlags; // SNTP flags: LI Leap Indicator, VN Version Number and Mode. + uint8_t mStratum; // Packet Stratum. + uint8_t mPoll; // Maximum interval between successive messages, in log2 seconds. + uint8_t mPrecision; // The precision of the system clock, in log2 seconds. + uint32_t mRootDelay; // Total round-trip delay to the reference clock, in NTP short format. + uint32_t mRootDispersion; // Total dispersion to the reference clock. + uint32_t mReferenceId; // ID identifying the particular server or reference clock. + uint32_t mReferenceTimestampSeconds; // Time the system clock was last set or corrected (NTP format). + uint32_t mReferenceTimestampFraction; // Fraction part of above value. + uint32_t mOriginateTimestampSeconds; // Time at client when request departed for the server (NTP format). + uint32_t mOriginateTimestampFraction; // Fraction part of above value. + uint32_t mReceiveTimestampSeconds; // Time at server when request arrived from the client (NTP format). + uint32_t mReceiveTimestampFraction; // Fraction part of above value. + uint32_t mTransmitTimestampSeconds; // Time at server when the response left for the client (NTP format). + uint32_t mTransmitTimestampFraction; // Fraction part of above value. + } OT_TOOL_PACKED_END; + + class QueryMetadata + { + public: + Error AppendTo(Message &aMessage) const { return aMessage.Append(*this); } + void ReadFrom(const Message &aMessage) + { + IgnoreError(aMessage.Read(aMessage.GetLength() - sizeof(*this), *this)); + } + + void UpdateIn(Message &aMessage) const { aMessage.Write(aMessage.GetLength() - sizeof(*this), *this); } + + uint32_t mTransmitTimestamp; // Time at client when request departed for server + Callback mResponseHandler; // Response handler callback + TimeMilli mTransmissionTime; // Time when the timer should shoot for this message + Ip6::Address mSourceAddress; // Source IPv6 address + Ip6::Address mDestinationAddress; // Destination IPv6 address + uint16_t mDestinationPort; // Destination UDP port + uint8_t mRetransmissionCount; // Number of retransmissions + }; + Message *NewMessage(const Header &aHeader); Message *CopyAndEnqueueMessage(const Message &aMessage, const QueryMetadata &aQueryMetadata); void DequeueMessage(Message &aMessage);