From 54afce92abacf177d2f67affd1b846d94a403770 Mon Sep 17 00:00:00 2001 From: Abtin Keshavarzian Date: Mon, 17 Jun 2024 18:31:26 -0700 Subject: [PATCH] [udp] add template `SocketIn` class for easier socket usage (#10392) This commit adds a new template `Udp::SocketIn`, which defines a UDP socket for use within a specified `Owner` class. It includes a `HandleUdpReceive()` member method callback for processing received messages. This model, similar to `TimerMilliIn`, eliminates the need for each socket user to define boilerplate code for providing a static member function that simply casts and calls the member method. This is now handled by the template `SocketIn` class. The `Udp::Socket::Open()` method is also modified to no longer require the callback and context parameters. --- src/core/coap/coap.cpp | 8 ++--- src/core/coap/coap.hpp | 7 +++-- src/core/meshcop/joiner_router.cpp | 9 ++---- src/core/meshcop/joiner_router.hpp | 6 ++-- src/core/meshcop/secure_transport.cpp | 9 ++---- src/core/meshcop/secure_transport.hpp | 6 ++-- src/core/net/dhcp6_client.cpp | 9 ++---- src/core/net/dhcp6_client.hpp | 6 ++-- src/core/net/dhcp6_server.cpp | 9 ++---- src/core/net/dhcp6_server.hpp | 11 ++++--- src/core/net/dns_client.cpp | 9 +++--- src/core/net/dns_client.hpp | 7 +++-- src/core/net/dnssd_server.cpp | 9 ++---- src/core/net/dnssd_server.hpp | 13 ++++---- src/core/net/sntp_client.cpp | 9 ++---- src/core/net/sntp_client.hpp | 10 +++---- src/core/net/srp_client.cpp | 8 ++--- src/core/net/srp_client.hpp | 7 +++-- src/core/net/srp_server.cpp | 9 ++---- src/core/net/srp_server.hpp | 4 +-- src/core/net/udp6.cpp | 17 +++++++++-- src/core/net/udp6.hpp | 43 +++++++++++++++++++++++---- src/core/thread/mle.cpp | 9 ++---- src/core/thread/mle.hpp | 20 ++++++------- tests/unit/test_srp_server.cpp | 4 +-- 25 files changed, 129 insertions(+), 129 deletions(-) diff --git a/src/core/coap/coap.cpp b/src/core/coap/coap.cpp index 941f6086b..38ba54e9b 100644 --- a/src/core/coap/coap.cpp +++ b/src/core/coap/coap.cpp @@ -1693,7 +1693,7 @@ Resource::Resource(Uri aUri, RequestHandler aHandler, void *aContext) Coap::Coap(Instance &aInstance) : CoapBase(aInstance, &Coap::Send) - , mSocket(aInstance) + , mSocket(aInstance, *this) { } @@ -1704,7 +1704,7 @@ Error Coap::Start(uint16_t aPort, Ip6::NetifIdentifier aNetifIdentifier) VerifyOrExit(!mSocket.IsBound()); - SuccessOrExit(error = mSocket.Open(&Coap::HandleUdpReceive, this)); + SuccessOrExit(error = mSocket.Open()); socketOpened = true; SuccessOrExit(error = mSocket.Bind(aPort, aNetifIdentifier)); @@ -1731,9 +1731,9 @@ exit: return error; } -void Coap::HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo) +void Coap::HandleUdpReceive(ot::Message &aMessage, const Ip6::MessageInfo &aMessageInfo) { - static_cast(aContext)->Receive(AsCoapMessage(aMessage), AsCoreType(aMessageInfo)); + Receive(AsCoapMessage(&aMessage), aMessageInfo); } Error Coap::Send(CoapBase &aCoapBase, ot::Message &aMessage, const Ip6::MessageInfo &aMessageInfo) diff --git a/src/core/coap/coap.hpp b/src/core/coap/coap.hpp index 9a08916a1..aeae3bb74 100644 --- a/src/core/coap/coap.hpp +++ b/src/core/coap/coap.hpp @@ -997,11 +997,14 @@ public: Error Stop(void); protected: - Ip6::Udp::Socket mSocket; + void HandleUdpReceive(ot::Message &aMessage, const Ip6::MessageInfo &aMessageInfo); + + using CoapSocket = Ip6::Udp::SocketIn; + + CoapSocket mSocket; private: static Error Send(CoapBase &aCoapBase, ot::Message &aMessage, const Ip6::MessageInfo &aMessageInfo); - static void HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo); Error Send(ot::Message &aMessage, const Ip6::MessageInfo &aMessageInfo); }; diff --git a/src/core/meshcop/joiner_router.cpp b/src/core/meshcop/joiner_router.cpp index a5df41367..b9f4a57c4 100644 --- a/src/core/meshcop/joiner_router.cpp +++ b/src/core/meshcop/joiner_router.cpp @@ -56,7 +56,7 @@ RegisterLogModule("JoinerRouter"); JoinerRouter::JoinerRouter(Instance &aInstance) : InstanceLocator(aInstance) - , mSocket(aInstance) + , mSocket(aInstance, *this) , mTimer(aInstance) , mJoinerUdpPort(0) , mIsJoinerPortConfigured(false) @@ -81,7 +81,7 @@ void JoinerRouter::Start(void) VerifyOrExit(!mSocket.IsBound()); - IgnoreError(mSocket.Open(&JoinerRouter::HandleUdpReceive, this)); + IgnoreError(mSocket.Open()); IgnoreError(mSocket.Bind(port)); IgnoreError(Get().AddUnsecurePort(port)); LogInfo("Joiner Router: start"); @@ -126,11 +126,6 @@ void JoinerRouter::SetJoinerUdpPort(uint16_t aJoinerUdpPort) Start(); } -void JoinerRouter::HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo) -{ - static_cast(aContext)->HandleUdpReceive(AsCoreType(aMessage), AsCoreType(aMessageInfo)); -} - void JoinerRouter::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo) { Error error; diff --git a/src/core/meshcop/joiner_router.hpp b/src/core/meshcop/joiner_router.hpp index bc76297ba..63672e2a5 100644 --- a/src/core/meshcop/joiner_router.hpp +++ b/src/core/meshcop/joiner_router.hpp @@ -100,8 +100,7 @@ private: void HandleNotifierEvents(Events aEvents); - static void HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo); - void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo); + void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo); template void HandleTmf(Coap::Message &aMessage, const Ip6::MessageInfo &aMessageInfo); @@ -120,8 +119,9 @@ private: Coap::Message *PrepareJoinerEntrustMessage(void); using JoinerRouterTimer = TimerMilliIn; + using JoinerSocket = Ip6::Udp::SocketIn; - Ip6::Udp::Socket mSocket; + JoinerSocket mSocket; JoinerRouterTimer mTimer; MessageQueue mDelayedJoinEnts; diff --git a/src/core/meshcop/secure_transport.cpp b/src/core/meshcop/secure_transport.cpp index e527c1f1e..bfb5db6d9 100644 --- a/src/core/meshcop/secure_transport.cpp +++ b/src/core/meshcop/secure_transport.cpp @@ -87,7 +87,7 @@ SecureTransport::SecureTransport(Instance &aInstance, bool aLayerTwoSecurity, bo , mMaxConnectionAttempts(0) , mRemainingConnectionAttempts(0) , mReceiveMessage(nullptr) - , mSocket(aInstance) + , mSocket(aInstance, *this) , mMessageSubType(Message::kSubTypeNone) , mMessageDefaultSubType(Message::kSubTypeNone) { @@ -158,7 +158,7 @@ Error SecureTransport::Open(ReceiveHandler aReceiveHandler, ConnectedHandler aCo VerifyOrExit(IsStateClosed(), error = kErrorAlready); - SuccessOrExit(error = mSocket.Open(&SecureTransport::HandleReceive, this)); + SuccessOrExit(error = mSocket.Open()); mConnectedCallback.Set(aConnectedHandler, aContext); mReceiveCallback.Set(aReceiveHandler, aContext); @@ -204,11 +204,6 @@ exit: return error; } -void SecureTransport::HandleReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo) -{ - static_cast(aContext)->HandleReceive(AsCoreType(aMessage), AsCoreType(aMessageInfo)); -} - void SecureTransport::HandleReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo) { VerifyOrExit(!IsStateClosed()); diff --git a/src/core/meshcop/secure_transport.hpp b/src/core/meshcop/secure_transport.hpp index 4641ef537..b5749a951 100644 --- a/src/core/meshcop/secure_transport.hpp +++ b/src/core/meshcop/secure_transport.hpp @@ -584,8 +584,6 @@ private: static void HandleTimer(Timer &aTimer); void HandleTimer(void); - static void HandleReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo); - void HandleReceive(const uint8_t *aBuf, uint16_t aLength); Error HandleSecureTransportSend(const uint8_t *aBuf, uint16_t aLength, Message::SubType aMessageSubType); @@ -595,6 +593,8 @@ private: static const char *StateToString(State aState); #endif + using TransportSocket = Ip6::Udp::SocketIn; + State mState; int mCipherSuites[2]; @@ -663,7 +663,7 @@ private: void *mContext; Ip6::MessageInfo mMessageInfo; - Ip6::Udp::Socket mSocket; + TransportSocket mSocket; Callback mTransportCallback; void *mTransportContext; diff --git a/src/core/net/dhcp6_client.cpp b/src/core/net/dhcp6_client.cpp index 11a88fd99..abf6e3265 100644 --- a/src/core/net/dhcp6_client.cpp +++ b/src/core/net/dhcp6_client.cpp @@ -52,7 +52,7 @@ RegisterLogModule("Dhcp6Client"); Client::Client(Instance &aInstance) : InstanceLocator(aInstance) - , mSocket(aInstance) + , mSocket(aInstance, *this) , mTrickleTimer(aInstance, Client::HandleTrickleTimer) , mStartTime(0) , mIdentityAssociationCurrent(nullptr) @@ -170,7 +170,7 @@ void Client::Start(void) { VerifyOrExit(!mSocket.IsBound()); - IgnoreError(mSocket.Open(&Client::HandleUdpReceive, this)); + IgnoreError(mSocket.Open()); IgnoreError(mSocket.Bind(kDhcpClientPort)); ProcessNextIdentityAssociation(); @@ -395,11 +395,6 @@ Error Client::AppendRapidCommit(Message &aMessage) return aMessage.Append(option); } -void Client::HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo) -{ - static_cast(aContext)->HandleUdpReceive(AsCoreType(aMessage), AsCoreType(aMessageInfo)); -} - void Client::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo) { OT_UNUSED_VARIABLE(aMessageInfo); diff --git a/src/core/net/dhcp6_client.hpp b/src/core/net/dhcp6_client.hpp index c50b13889..2d9287db0 100644 --- a/src/core/net/dhcp6_client.hpp +++ b/src/core/net/dhcp6_client.hpp @@ -126,8 +126,7 @@ private: Error AppendElapsedTime(Message &aMessage); Error AppendRapidCommit(Message &aMessage); - static void HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo); - void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo); + void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo); void ProcessReply(Message &aMessage); uint16_t FindOption(Message &aMessage, uint16_t aOffset, uint16_t aLength, Code aCode); @@ -140,8 +139,9 @@ private: static void HandleTrickleTimer(TrickleTimer &aTrickleTimer); void HandleTrickleTimer(void); - Ip6::Udp::Socket mSocket; + using ClientSocket = Ip6::Udp::SocketIn; + ClientSocket mSocket; TrickleTimer mTrickleTimer; TransactionId mTransactionId; diff --git a/src/core/net/dhcp6_server.cpp b/src/core/net/dhcp6_server.cpp index ea9570737..f25884e12 100644 --- a/src/core/net/dhcp6_server.cpp +++ b/src/core/net/dhcp6_server.cpp @@ -52,7 +52,7 @@ RegisterLogModule("Dhcp6Server"); Server::Server(Instance &aInstance) : InstanceLocator(aInstance) - , mSocket(aInstance) + , mSocket(aInstance, *this) , mPrefixAgentsCount(0) , mPrefixAgentsMask(0) { @@ -138,7 +138,7 @@ void Server::Start(void) { VerifyOrExit(!mSocket.IsOpen()); - IgnoreError(mSocket.Open(&Server::HandleUdpReceive, this)); + IgnoreError(mSocket.Open()); IgnoreError(mSocket.Bind(kDhcpServerPort)); exit: @@ -176,11 +176,6 @@ exit: OT_UNUSED_VARIABLE(error); } -void Server::HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo) -{ - static_cast(aContext)->HandleUdpReceive(AsCoreType(aMessage), AsCoreType(aMessageInfo)); -} - void Server::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo) { Header header; diff --git a/src/core/net/dhcp6_server.hpp b/src/core/net/dhcp6_server.hpp index eba2d6c96..8869f24e6 100644 --- a/src/core/net/dhcp6_server.hpp +++ b/src/core/net/dhcp6_server.hpp @@ -192,11 +192,8 @@ private: Error AppendVendorSpecificInformation(Message &aMessage); Error AddIaAddress(Message &aMessage, const Ip6::Address &aPrefix, ClientIdentifier &aClientId); - - static void HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo); - void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo); - - void ProcessSolicit(Message &aMessage, const Ip6::Address &aDst, const TransactionId &aTransactionId); + void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo); + void ProcessSolicit(Message &aMessage, const Ip6::Address &aDst, const TransactionId &aTransactionId); uint16_t FindOption(Message &aMessage, uint16_t aOffset, uint16_t aLength, Code aCode); Error ProcessClientIdentifier(Message &aMessage, uint16_t aOffset, ClientIdentifier &aClientId); @@ -209,7 +206,9 @@ private: ClientIdentifier &aClientId, IaNa &aIaNa); - Ip6::Udp::Socket mSocket; + using ServerSocket = Ip6::Udp::SocketIn; + + ServerSocket mSocket; PrefixAgent mPrefixAgents[kNumPrefixes]; uint8_t mPrefixAgentsCount; diff --git a/src/core/net/dns_client.cpp b/src/core/net/dns_client.cpp index 61bdb1714..1b1b77ee5 100644 --- a/src/core/net/dns_client.cpp +++ b/src/core/net/dns_client.cpp @@ -737,7 +737,7 @@ const uint16_t *const Client::kQuestionRecordTypes[] = { Client::Client(Instance &aInstance) : InstanceLocator(aInstance) - , mSocket(aInstance) + , mSocket(aInstance, *this) #if OPENTHREAD_CONFIG_DNS_CLIENT_OVER_TCP_ENABLE , mTcpState(kTcpUninitialized) #endif @@ -768,7 +768,7 @@ Error Client::Start(void) { Error error; - SuccessOrExit(error = mSocket.Open(&Client::HandleUdpReceive, this)); + SuccessOrExit(error = mSocket.Open()); SuccessOrExit(error = mSocket.Bind(0, Ip6::kNetifUnspecified)); exit: @@ -1301,11 +1301,10 @@ exit: return matchedQuery; } -void Client::HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMsgInfo) +void Client::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMsgInfo) { OT_UNUSED_VARIABLE(aMsgInfo); - - static_cast(aContext)->ProcessResponse(AsCoreType(aMessage)); + ProcessResponse(aMessage); } void Client::ProcessResponse(const Message &aResponseMessage) diff --git a/src/core/net/dns_client.hpp b/src/core/net/dns_client.hpp index e086f746b..9ac6a3868 100644 --- a/src/core/net/dns_client.hpp +++ b/src/core/net/dns_client.hpp @@ -839,7 +839,7 @@ private: static void GetQueryTypeAndCallback(const Query &aQuery, QueryType &aType, Callback &aCallback, void *&aContext); Error AppendNameFromQuery(const Query &aQuery, Message &aMessage); Query *FindQueryById(uint16_t aMessageId); - static void HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMsgInfo); + void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMsgInfo); void ProcessResponse(const Message &aResponseMessage); Error ParseResponse(const Message &aResponseMessage, Query *&aQuery, Error &aResponseError); bool CanFinalizeQuery(Query &aQuery); @@ -904,9 +904,10 @@ private: static constexpr uint16_t kUdpQueryMaxSize = 512; - using RetryTimer = TimerMilliIn; + using RetryTimer = TimerMilliIn; + using ClientSocket = Ip6::Udp::SocketIn; - Ip6::Udp::Socket mSocket; + ClientSocket mSocket; #if OPENTHREAD_CONFIG_DNS_CLIENT_OVER_TCP_ENABLE Ip6::Tcp::Endpoint mEndpoint; diff --git a/src/core/net/dnssd_server.cpp b/src/core/net/dnssd_server.cpp index 9f177db4b..10fa87d9f 100644 --- a/src/core/net/dnssd_server.cpp +++ b/src/core/net/dnssd_server.cpp @@ -63,7 +63,7 @@ const char *Server::kBlockedDomains[] = {"ipv4only.arpa."}; Server::Server(Instance &aInstance) : InstanceLocator(aInstance) - , mSocket(aInstance) + , mSocket(aInstance, *this) #if OPENTHREAD_CONFIG_DNSSD_DISCOVERY_PROXY_ENABLE , mDiscoveryProxy(aInstance) #endif @@ -82,7 +82,7 @@ Error Server::Start(void) VerifyOrExit(!IsRunning()); - SuccessOrExit(error = mSocket.Open(&Server::HandleUdpReceive, this)); + SuccessOrExit(error = mSocket.Open()); SuccessOrExit(error = mSocket.Bind(kPort, kBindUnspecifiedNetif ? Ip6::kNetifUnspecified : Ip6::kNetifThread)); #if OPENTHREAD_CONFIG_SRP_SERVER_ENABLE @@ -135,11 +135,6 @@ void Server::Stop(void) #endif } -void Server::HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo) -{ - static_cast(aContext)->HandleUdpReceive(AsCoreType(aMessage), AsCoreType(aMessageInfo)); -} - void Server::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo) { Request request; diff --git a/src/core/net/dnssd_server.hpp b/src/core/net/dnssd_server.hpp index 692239ecb..ea8f7e7d7 100644 --- a/src/core/net/dnssd_server.hpp +++ b/src/core/net/dnssd_server.hpp @@ -513,11 +513,9 @@ private: }; #endif - bool IsRunning(void) const { return mSocket.IsBound(); } - static void HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo); - void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo); - void ProcessQuery(Request &aRequest); - + bool IsRunning(void) const { return mSocket.IsBound(); } + void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo); + void ProcessQuery(Request &aRequest); void ResolveByProxy(Response &aResponse, const Ip6::MessageInfo &aMessageInfo); void RemoveQueryAndPrepareResponse(ProxyQuery &aQuery, ProxyQueryInfo &aInfo, Response &aResponse); void Finalize(ProxyQuery &aQuery, ResponseCode aResponseCode); @@ -560,7 +558,8 @@ private: void UpdateResponseCounters(ResponseCode aResponseCode); - using ServerTimer = TimerMilliIn; + using ServerTimer = TimerMilliIn; + using ServerSocket = Ip6::Udp::SocketIn; static const char kDefaultDomainName[]; static const char kSubLabel[]; @@ -568,7 +567,7 @@ private: static const char *kBlockedDomains[]; #endif - Ip6::Udp::Socket mSocket; + ServerSocket mSocket; ProxyQueryList mProxyQueries; Callback mQuerySubscribe; diff --git a/src/core/net/sntp_client.cpp b/src/core/net/sntp_client.cpp index 3950f49b1..7e28ad795 100644 --- a/src/core/net/sntp_client.cpp +++ b/src/core/net/sntp_client.cpp @@ -51,7 +51,7 @@ namespace Sntp { RegisterLogModule("SntpClnt"); Client::Client(Instance &aInstance) - : mSocket(aInstance) + : mSocket(aInstance, *this) , mRetransmissionTimer(aInstance) , mUnixEra(0) { @@ -61,7 +61,7 @@ Error Client::Start(void) { Error error; - SuccessOrExit(error = mSocket.Open(&Client::HandleUdpReceive, this)); + SuccessOrExit(error = mSocket.Open()); SuccessOrExit(error = mSocket.Bind(0, Ip6::kNetifUnspecified)); exit: @@ -263,11 +263,6 @@ void Client::HandleRetransmissionTimer(void) mRetransmissionTimer.FireAt(nextTime); } -void Client::HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo) -{ - static_cast(aContext)->HandleUdpReceive(AsCoreType(aMessage), AsCoreType(aMessageInfo)); -} - void Client::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo) { OT_UNUSED_VARIABLE(aMessageInfo); diff --git a/src/core/net/sntp_client.hpp b/src/core/net/sntp_client.hpp index d6d4dad3c..dd4c4ab33 100644 --- a/src/core/net/sntp_client.hpp +++ b/src/core/net/sntp_client.hpp @@ -272,14 +272,12 @@ private: void FinalizeSntpTransaction(Message &aQuery, const QueryMetadata &aQueryMetadata, uint64_t aTime, Error aResult); void HandleRetransmissionTimer(void); + void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo); - static void HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo); - void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo); - - using RetxTimer = TimerMilliIn; - - Ip6::Udp::Socket mSocket; + using RetxTimer = TimerMilliIn; + using ClientSocket = Ip6::Udp::SocketIn; + ClientSocket mSocket; MessageQueue mPendingQueries; RetxTimer mRetransmissionTimer; diff --git a/src/core/net/srp_client.cpp b/src/core/net/srp_client.cpp index 6b6a15250..9ac4e72db 100644 --- a/src/core/net/srp_client.cpp +++ b/src/core/net/srp_client.cpp @@ -267,7 +267,7 @@ Client::Client(Instance &aInstance) , mKeyLease(0) , mDefaultLease(kDefaultLease) , mDefaultKeyLease(kDefaultKeyLease) - , mSocket(aInstance) + , mSocket(aInstance, *this) , mDomainName(kDefaultDomainName) , mTimer(aInstance) { @@ -296,7 +296,7 @@ Error Client::Start(const Ip6::SockAddr &aServerSockAddr, Requester aRequester) VerifyOrExit(GetState() == kStateStopped, error = (aServerSockAddr == GetServerAddress()) ? kErrorNone : kErrorBusy); - SuccessOrExit(error = mSocket.Open(Client::HandleUdpReceive, this)); + SuccessOrExit(error = mSocket.Open()); SuccessOrExit(error = mSocket.Connect(aServerSockAddr)); LogInfo("%starting, server %s", (aRequester == kRequesterUser) ? "S" : "Auto-s", @@ -1583,11 +1583,11 @@ void Client::UpdateRecordLengthInMessage(Dns::ResourceRecord &aRecord, uint16_t aMessage.Write(aOffset, aRecord); } -void Client::HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo) +void Client::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo) { OT_UNUSED_VARIABLE(aMessageInfo); - static_cast(aContext)->ProcessResponse(AsCoreType(aMessage)); + ProcessResponse(aMessage); } void Client::ProcessResponse(Message &aMessage) diff --git a/src/core/net/srp_client.hpp b/src/core/net/srp_client.hpp index 8c9cf7e3f..f45b95364 100644 --- a/src/core/net/srp_client.hpp +++ b/src/core/net/srp_client.hpp @@ -1040,7 +1040,7 @@ private: Error AppendUpdateLeaseOptRecord(Message &aMessage); Error AppendSignature(Message &aMessage, Info &aInfo); void UpdateRecordLengthInMessage(Dns::ResourceRecord &aRecord, uint16_t aOffset, Message &aMessage) const; - static void HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo); + void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo); void ProcessResponse(Message &aMessage); bool IsResponseMessageIdValid(uint16_t aId) const; void HandleUpdateDone(void); @@ -1074,7 +1074,8 @@ private: static_assert(kMaxTxFailureRetries < 16, "kMaxTxFailureRetries exceed the range of mTxFailureRetryCount (4-bit)"); - using DelayTimer = TimerMilliIn; + using DelayTimer = TimerMilliIn; + using ClientSocket = Ip6::Udp::SocketIn; State mState; uint8_t mTxFailureRetryCount : 4; @@ -1097,7 +1098,7 @@ private: uint32_t mDefaultLease; uint32_t mDefaultKeyLease; - Ip6::Udp::Socket mSocket; + ClientSocket mSocket; Callback mCallback; const char *mDomainName; diff --git a/src/core/net/srp_server.cpp b/src/core/net/srp_server.cpp index b3ea87b5d..7fafba1d9 100644 --- a/src/core/net/srp_server.cpp +++ b/src/core/net/srp_server.cpp @@ -86,7 +86,7 @@ static Dns::UpdateHeader::Response ErrorToDnsResponseCode(Error aError) Server::Server(Instance &aInstance) : InstanceLocator(aInstance) - , mSocket(aInstance) + , mSocket(aInstance, *this) , mLeaseTimer(aInstance) , mOutstandingUpdatesTimer(aInstance) , mCompletedUpdateTask(aInstance) @@ -668,7 +668,7 @@ Error Server::PrepareSocket(void) #endif VerifyOrExit(!mSocket.IsOpen()); - SuccessOrExit(error = mSocket.Open(HandleUdpReceive, this)); + SuccessOrExit(error = mSocket.Open()); error = mSocket.Bind(mPort, Ip6::kNetifThread); exit: @@ -1561,11 +1561,6 @@ exit: FreeMessageOnError(response, error); } -void Server::HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo) -{ - static_cast(aContext)->HandleUdpReceive(AsCoreType(aMessage), AsCoreType(aMessageInfo)); -} - void Server::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo) { Error error = ProcessMessage(aMessage, aMessageInfo); diff --git a/src/core/net/srp_server.hpp b/src/core/net/srp_server.hpp index ebe67d8fe..48c9780a3 100644 --- a/src/core/net/srp_server.hpp +++ b/src/core/net/srp_server.hpp @@ -1034,7 +1034,6 @@ private: uint32_t aKeyLease, bool mUseShortLeaseOption, const Ip6::MessageInfo &aMessageInfo); - static void HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo); void HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo); void HandleLeaseTimer(void); static void HandleOutstandingUpdatesTimer(Timer &aTimer); @@ -1050,8 +1049,9 @@ private: using LeaseTimer = TimerMilliIn; using UpdateTimer = TimerMilliIn; using CompletedUpdatesTask = TaskletIn; + using ServerSocket = Ip6::Udp::SocketIn; - Ip6::Udp::Socket mSocket; + ServerSocket mSocket; Callback mServiceUpdateHandler; diff --git a/src/core/net/udp6.cpp b/src/core/net/udp6.cpp index 05ba27bc9..5fd34b266 100644 --- a/src/core/net/udp6.cpp +++ b/src/core/net/udp6.cpp @@ -47,6 +47,9 @@ namespace ot { namespace Ip6 { +//--------------------------------------------------------------------------------------------------------------------- +// Udp::SocketHandle + bool Udp::SocketHandle::Matches(const MessageInfo &aMessageInfo) const { bool matches = false; @@ -71,10 +74,15 @@ exit: return matches; } -Udp::Socket::Socket(Instance &aInstance) +//--------------------------------------------------------------------------------------------------------------------- +// Udp::Socket + +Udp::Socket::Socket(Instance &aInstance, ReceiveHandler aHandler, void *aContext) : InstanceLocator(aInstance) { Clear(); + mHandler = aHandler; + mContext = aContext; } Message *Udp::Socket::NewMessage(void) { return NewMessage(0); } @@ -86,7 +94,7 @@ Message *Udp::Socket::NewMessage(uint16_t aReserved, const Message::Settings &aS return Get().NewMessage(aReserved, aSettings); } -Error Udp::Socket::Open(otUdpReceive aHandler, void *aContext) { return Get().Open(*this, aHandler, aContext); } +Error Udp::Socket::Open(void) { return Get().Open(*this, mHandler, mContext); } bool Udp::Socket::IsOpen(void) const { return Get().IsOpen(*this); } @@ -147,6 +155,9 @@ exit: } #endif +//--------------------------------------------------------------------------------------------------------------------- +// Udp + Udp::Udp(Instance &aInstance) : InstanceLocator(aInstance) , mEphemeralPort(kDynamicPortMin) @@ -169,7 +180,7 @@ exit: return error; } -Error Udp::Open(SocketHandle &aSocket, otUdpReceive aHandler, void *aContext) +Error Udp::Open(SocketHandle &aSocket, ReceiveHandler aHandler, void *aContext) { Error error = kErrorNone; diff --git a/src/core/net/udp6.hpp b/src/core/net/udp6.hpp index 24e706607..73d7429fc 100644 --- a/src/core/net/udp6.hpp +++ b/src/core/net/udp6.hpp @@ -84,6 +84,8 @@ enum NetifIdentifier : uint8_t class Udp : public InstanceLocator, private NonCopyable { public: + typedef otUdpReceive ReceiveHandler; ///< Receive handler callback. + /** * Implements a UDP/IPv6 socket. * @@ -157,9 +159,11 @@ public: * Initializes the object. * * @param[in] aInstance A reference to OpenThread instance. + * @param[in] aHandler A pointer to a function that is called when receiving UDP messages. + * @param[in] aContext A pointer to arbitrary context information. * */ - explicit Socket(Instance &aInstance); + Socket(Instance &aInstance, ReceiveHandler aHandler, void *aContext); /** * Returns a new UDP message with default settings (link security enabled and `kPriorityNormal`) @@ -193,14 +197,11 @@ public: /** * Opens the UDP socket. * - * @param[in] aHandler A pointer to a function that is called when receiving UDP messages. - * @param[in] aContext A pointer to arbitrary context information. - * * @retval kErrorNone Successfully opened the socket. * @retval kErrorFailed Failed to open the socket. * */ - Error Open(otUdpReceive aHandler, void *aContext); + Error Open(void); /** * Returns if the UDP socket is open. @@ -324,6 +325,36 @@ public: #endif }; + /** + * A socket owned by a specific type with a given owner type as the callback. + * + * @tparam Owner The type of the owner of this socket. + * @tparam HandleUdpReceivePtr A pointer to a non-static member method of `Owner` to handle received messages. + * + */ + template + class SocketIn : public Socket + { + public: + /** + * Initializes the socket. + * + * @param[in] aInstance The OpenThread instance. + * @param[in] aOnwer The owner of the socket, providing the `HandleUdpReceivePtr` callback. + * + */ + explicit SocketIn(Instance &aInstance, Owner &aOwner) + : Socket(aInstance, HandleUdpReceive, &aOwner) + { + } + + private: + static void HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo) + { + (reinterpret_cast(aContext)->*HandleUdpReceivePtr)(AsCoreType(aMessage), AsCoreType(aMessageInfo)); + } + }; + /** * Implements a UDP receiver. * @@ -480,7 +511,7 @@ public: * @retval kErrorFailed Failed to open the socket. * */ - Error Open(SocketHandle &aSocket, otUdpReceive aHandler, void *aContext); + Error Open(SocketHandle &aSocket, ReceiveHandler aHandler, void *aContext); /** * Returns if a UDP socket is open. diff --git a/src/core/thread/mle.cpp b/src/core/thread/mle.cpp index 57b1beff9..7b78baf74 100644 --- a/src/core/thread/mle.cpp +++ b/src/core/thread/mle.cpp @@ -106,7 +106,7 @@ Mle::Mle(Instance &aInstance) #endif , mAlternateTimestamp(0) , mNeighborTable(aInstance) - , mSocket(aInstance) + , mSocket(aInstance, *this) #if OPENTHREAD_CONFIG_PARENT_SEARCH_ENABLE , mParentSearch(aInstance) #endif @@ -150,7 +150,7 @@ Error Mle::Enable(void) Error error = kErrorNone; UpdateLinkLocalAddress(); - SuccessOrExit(error = mSocket.Open(&Mle::HandleUdpReceive, this)); + SuccessOrExit(error = mSocket.Open()); SuccessOrExit(error = mSocket.Bind(kUdpPort)); #if OPENTHREAD_CONFIG_PARENT_SEARCH_ENABLE @@ -2390,11 +2390,6 @@ exit: return error; } -void Mle::HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo) -{ - static_cast(aContext)->HandleUdpReceive(AsCoreType(aMessage), AsCoreType(aMessageInfo)); -} - void Mle::HandleUdpReceive(Message &aMessage, const Ip6::MessageInfo &aMessageInfo) { Error error = kErrorNone; diff --git a/src/core/thread/mle.hpp b/src/core/thread/mle.hpp index d920e95fa..85d1207cc 100644 --- a/src/core/thread/mle.hpp +++ b/src/core/thread/mle.hpp @@ -1222,9 +1222,6 @@ private: //------------------------------------------------------------------------------------------------------------------ // Methods - static void HandleUdpReceive(void *aContext, otMessage *aMessage, const otMessageInfo *aMessageInfo); - static void HandleDetachGracefullyTimer(Timer &aTimer); - Error Start(StartMode aMode); void Stop(StopMode aMode); TxMessage *NewMleMessage(Command aCommand); @@ -1364,6 +1361,7 @@ private: using AttachTimer = TimerMilliIn; using DelayTimer = TimerMilliIn; using MsgTxTimer = TimerMilliIn; + using MleSocket = Ip6::Udp::SocketIn; static const otMeshLocalPrefix kMeshLocalPrefixInit; @@ -1407,14 +1405,14 @@ private: uint64_t mLastUpdatedTimestamp; #endif - LeaderData mLeaderData; - Parent mParent; - NeighborTable mNeighborTable; - MessageQueue mDelayedResponses; - TxChallenge mParentRequestChallenge; - ParentCandidate mParentCandidate; - Ip6::Udp::Socket mSocket; - Counters mCounters; + LeaderData mLeaderData; + Parent mParent; + NeighborTable mNeighborTable; + MessageQueue mDelayedResponses; + TxChallenge mParentRequestChallenge; + ParentCandidate mParentCandidate; + MleSocket mSocket; + Counters mCounters; #if OPENTHREAD_CONFIG_PARENT_SEARCH_ENABLE ParentSearch mParentSearch; #endif diff --git a/tests/unit/test_srp_server.cpp b/tests/unit/test_srp_server.cpp index 15bcdc104..f0218e004 100644 --- a/tests/unit/test_srp_server.cpp +++ b/tests/unit/test_srp_server.cpp @@ -1068,7 +1068,7 @@ void TestSrpClientDelayedResponse(void) //- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - // Prepare a socket to act as SRP server. - Ip6::Udp::Socket udpSocket(*sInstance); + Ip6::Udp::Socket udpSocket(*sInstance, HandleServerUdpReceive, nullptr); Ip6::SockAddr serverSockAddr; uint16_t firstMsgId; Message *response; @@ -1076,7 +1076,7 @@ void TestSrpClientDelayedResponse(void) sServerRxCount = 0; - SuccessOrQuit(udpSocket.Open(HandleServerUdpReceive, nullptr)); + SuccessOrQuit(udpSocket.Open()); SuccessOrQuit(udpSocket.Bind(kServerPort, Ip6::kNetifThread)); //- - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -