diff --git a/cpp/core/BUILD b/cpp/core/BUILD index 2bbf7be8..f0b551ee 100644 --- a/cpp/core/BUILD +++ b/cpp/core/BUILD @@ -21,7 +21,7 @@ cc_library( "core.h", ], visibility = [ - "//googlemac/iPhone/Shared/Nearby/Connections:__subpackages__", + "//googlemac/iPhone/Shared/Nearby/Connections_v2:__subpackages__", ], deps = [ ":core_types", diff --git a/cpp/core/internal/base_endpoint_channel.cc b/cpp/core/internal/base_endpoint_channel.cc index 8a97dee3..30e8c8be 100644 --- a/cpp/core/internal/base_endpoint_channel.cc +++ b/cpp/core/internal/base_endpoint_channel.cc @@ -259,6 +259,11 @@ std::string BaseEndpointChannel::GetType() const { std::string BaseEndpointChannel::GetName() const { return channel_name_; } +int BaseEndpointChannel::GetMaxTransmitPacketSize() const { + // Return default value if the medium never define it's chunk size. + return kDefaultMaxTransmitPacketSize; +} + void BaseEndpointChannel::EnableEncryption( std::shared_ptr context) { MutexLock crypto_lock(&crypto_mutex_); diff --git a/cpp/core/internal/base_endpoint_channel.h b/cpp/core/internal/base_endpoint_channel.h index d72309d4..5b64170c 100644 --- a/cpp/core/internal/base_endpoint_channel.h +++ b/cpp/core/internal/base_endpoint_channel.h @@ -62,6 +62,10 @@ class BaseEndpointChannel : public EndpointChannel { // Returns the name of the EndpointChannel. std::string GetName() const override; + // Returns the maximum supported transmit packet size(MTU) for the underlying + // transport. + int GetMaxTransmitPacketSize() const override; + // Enables encryption on the EndpointChannel. // Should be called after connection is accepted by both parties, and // before entering data phase, where Payloads may be exchanged. @@ -92,6 +96,9 @@ class BaseEndpointChannel : public EndpointChannel { // Used to sanity check that our frame sizes are reasonable. static constexpr std::int32_t kMaxAllowedReadBytes = 1048576; // 1MB + // The default maximum transmit unit/packet size. + static constexpr int kDefaultMaxTransmitPacketSize = 65536; // 64 KB + bool IsEncryptionEnabledLocked() const ABSL_EXCLUSIVE_LOCKS_REQUIRED(crypto_mutex_); void UnblockPausedWriter() ABSL_EXCLUSIVE_LOCKS_REQUIRED(is_paused_mutex_); diff --git a/cpp/core/internal/ble_endpoint_channel.cc b/cpp/core/internal/ble_endpoint_channel.cc index 181327ce..5c8d479a 100644 --- a/cpp/core/internal/ble_endpoint_channel.cc +++ b/cpp/core/internal/ble_endpoint_channel.cc @@ -47,6 +47,10 @@ proto::connections::Medium BleEndpointChannel::GetMedium() const { return proto::connections::Medium::BLE; } +int BleEndpointChannel::GetMaxTransmitPacketSize() const { + return kDefaultBleMaxTransmitPacketSize; +} + void BleEndpointChannel::CloseImpl() { auto status = ble_socket_.Close(); if (!status.Ok()) { diff --git a/cpp/core/internal/ble_endpoint_channel.h b/cpp/core/internal/ble_endpoint_channel.h index b7ecc5aa..3d358903 100644 --- a/cpp/core/internal/ble_endpoint_channel.h +++ b/cpp/core/internal/ble_endpoint_channel.h @@ -30,7 +30,11 @@ class BleEndpointChannel final : public BaseEndpointChannel { proto::connections::Medium GetMedium() const override; + int GetMaxTransmitPacketSize() const override; + private: + static constexpr int kDefaultBleMaxTransmitPacketSize = 512; // 512 bytes + void CloseImpl() override; BleSocket ble_socket_; diff --git a/cpp/core/internal/bluetooth_endpoint_channel.cc b/cpp/core/internal/bluetooth_endpoint_channel.cc index a1a9ecdc..56939081 100644 --- a/cpp/core/internal/bluetooth_endpoint_channel.cc +++ b/cpp/core/internal/bluetooth_endpoint_channel.cc @@ -47,6 +47,10 @@ proto::connections::Medium BluetoothEndpointChannel::GetMedium() const { return proto::connections::Medium::BLUETOOTH; } +int BluetoothEndpointChannel::GetMaxTransmitPacketSize() const { + return kDefaultBTMaxTransmitPacketSize; +} + void BluetoothEndpointChannel::CloseImpl() { auto status = bluetooth_socket_.Close(); if (!status.Ok()) { diff --git a/cpp/core/internal/bluetooth_endpoint_channel.h b/cpp/core/internal/bluetooth_endpoint_channel.h index 6f9d9757..06b91957 100644 --- a/cpp/core/internal/bluetooth_endpoint_channel.h +++ b/cpp/core/internal/bluetooth_endpoint_channel.h @@ -33,7 +33,11 @@ class BluetoothEndpointChannel final : public BaseEndpointChannel { proto::connections::Medium GetMedium() const override; + int GetMaxTransmitPacketSize() const override; + private: + static constexpr int kDefaultBTMaxTransmitPacketSize = 1980; // 990 * 2 Bytes + void CloseImpl() override; BluetoothSocket bluetooth_socket_; diff --git a/cpp/core/internal/encryption_runner_test.cc b/cpp/core/internal/encryption_runner_test.cc index de0d97b1..3fdc0180 100644 --- a/cpp/core/internal/encryption_runner_test.cc +++ b/cpp/core/internal/encryption_runner_test.cc @@ -54,6 +54,7 @@ class FakeEndpointChannel : public EndpointChannel { std::string GetType() const override { return "fake-channel-type"; } std::string GetName() const override { return "fake-channel"; } Medium GetMedium() const override { return Medium::BLE; } + int GetMaxTransmitPacketSize() const override { return 512; } void EnableEncryption(std::shared_ptr context) override {} void DisableEncryption() override {} bool IsPaused() const override { return false; } diff --git a/cpp/core/internal/endpoint_channel.h b/cpp/core/internal/endpoint_channel.h index 39cf7309..02b9c0d9 100644 --- a/cpp/core/internal/endpoint_channel.h +++ b/cpp/core/internal/endpoint_channel.h @@ -56,6 +56,10 @@ class EndpointChannel { // Returns the analytics enum representing the medium of this EndpointChannel. virtual proto::connections::Medium GetMedium() const = 0; + // Returns the maximum supported transmit packet size(MTU) for the underlying + // transport. + virtual int GetMaxTransmitPacketSize() const = 0; + // Enables encryption on the EndpointChannel. virtual void EnableEncryption(std::shared_ptr context) = 0; diff --git a/cpp/core/internal/endpoint_manager.cc b/cpp/core/internal/endpoint_manager.cc index ba379304..3c080e3a 100644 --- a/cpp/core/internal/endpoint_manager.cc +++ b/cpp/core/internal/endpoint_manager.cc @@ -417,16 +417,14 @@ void EndpointManager::UnregisterEndpoint(ClientProxy* client, latch.Await(); } -// Designed to run asynchronously. It is called from IO thread pools, and -// jobs in these pools may be waited for from the EndpointManager thread. If we -// allow synchronous behavior here it will cause a live lock. -void EndpointManager::DiscardEndpoint(ClientProxy* client, - const std::string& endpoint_id) { - RunOnEndpointManagerThread([this, client, endpoint_id]() { - RemoveEndpoint(client, endpoint_id, - /*notify=*/ - client->IsConnectedToEndpoint(endpoint_id)); - }); +int EndpointManager::GetMaxTransmitPacketSize(const std::string& endpoint_id) { + std::shared_ptr channel = + channel_manager_->GetChannelForEndpoint(endpoint_id); + if (channel == nullptr) { + return 0; + } + + return channel->GetMaxTransmitPacketSize(); } std::vector EndpointManager::SendPayloadChunk( @@ -441,6 +439,18 @@ std::vector EndpointManager::SendPayloadChunk( /*packet_type=*/"DATA"); } +// Designed to run asynchronously. It is called from IO thread pools, and +// jobs in these pools may be waited for from the EndpointManager thread. If we +// allow synchronous behavior here it will cause a live lock. +void EndpointManager::DiscardEndpoint(ClientProxy* client, + const std::string& endpoint_id) { + RunOnEndpointManagerThread([this, client, endpoint_id]() { + RemoveEndpoint(client, endpoint_id, + /*notify=*/ + client->IsConnectedToEndpoint(endpoint_id)); + }); +} + std::vector EndpointManager::SendControlMessage( const PayloadTransferFrame::PayloadHeader& header, const PayloadTransferFrame::ControlMessage& control, diff --git a/cpp/core/internal/endpoint_manager.h b/cpp/core/internal/endpoint_manager.h index 6028b6e5..476b0c96 100644 --- a/cpp/core/internal/endpoint_manager.h +++ b/cpp/core/internal/endpoint_manager.h @@ -112,6 +112,10 @@ class EndpointManager { // this case, we do not notify the client of onDisconnected(). void UnregisterEndpoint(ClientProxy* client, const std::string& endpoint_id); + // Returns the maximum supported transmit packet size(MTU) for the underlying + // transport. + int GetMaxTransmitPacketSize(const std::string& endpoint_id); + // Returns the list of endpoints to which sending this chunk failed. // // Invoked from the PayloadManager's sendPayload() method. diff --git a/cpp/core/internal/endpoint_manager_test.cc b/cpp/core/internal/endpoint_manager_test.cc index 4cd64752..a43e0766 100644 --- a/cpp/core/internal/endpoint_manager_test.cc +++ b/cpp/core/internal/endpoint_manager_test.cc @@ -54,6 +54,7 @@ class MockEndpointChannel : public EndpointChannel { MOCK_METHOD(std::string, GetType, (), (const override)); MOCK_METHOD(std::string, GetName, (), (const override)); MOCK_METHOD(Medium, GetMedium, (), (const override)); + MOCK_METHOD(int, GetMaxTransmitPacketSize, (), (const override)); MOCK_METHOD(void, EnableEncryption, (std::shared_ptr context), (override)); MOCK_METHOD(void, DisableEncryption, (), (override)); diff --git a/cpp/core/internal/internal_payload.h b/cpp/core/internal/internal_payload.h index 7b1cb699..8861f17e 100644 --- a/cpp/core/internal/internal_payload.h +++ b/cpp/core/internal/internal_payload.h @@ -63,8 +63,10 @@ class InternalPayload { // byte blobs for sending across a hard boundary (like the other side of // a Binder, or another device altogether). // + // @param chunk_size The preferred size of the next chunk. Depending on + // payload type, the provided size may be ignored. // @return The next chunk from the Payload, or null if we've reached the end. - virtual ByteArray DetachNextChunk() = 0; + virtual ByteArray DetachNextChunk(int chunk_size) = 0; // Adds the next chunk that comprises the Payload to which this object is // bound. diff --git a/cpp/core/internal/internal_payload_factory.cc b/cpp/core/internal/internal_payload_factory.cc index 08e30dd5..070bd513 100644 --- a/cpp/core/internal/internal_payload_factory.cc +++ b/cpp/core/internal/internal_payload_factory.cc @@ -22,6 +22,7 @@ #include "platform/base/exception.h" #include "platform/public/condition_variable.h" #include "platform/public/file.h" +#include "platform/public/logging.h" #include "platform/public/mutex.h" #include "platform/public/pipe.h" #include "absl/memory/memory.h" @@ -47,7 +48,7 @@ class BytesInternalPayload : public InternalPayload { // Relinquishes ownership of the payload_; retrieves and returns the stored // ByteArray. - ByteArray DetachNextChunk() override { + ByteArray DetachNextChunk(int chunk_size) override { if (detached_only_chunk_) { return {}; } @@ -80,11 +81,11 @@ class OutgoingStreamInternalPayload : public InternalPayload { std::int64_t GetTotalSize() const override { return -1; } - ByteArray DetachNextChunk() override { + ByteArray DetachNextChunk(int chunk_size) override { InputStream* input_stream = payload_.AsStream(); if (!input_stream) return {}; - ExceptionOr bytes_read = input_stream->Read(kChunkSize); + ExceptionOr bytes_read = input_stream->Read(chunk_size); if (!bytes_read.ok()) { input_stream->Close(); return {}; @@ -93,8 +94,8 @@ class OutgoingStreamInternalPayload : public InternalPayload { ByteArray scoped_bytes_read = std::move(bytes_read.result()); if (scoped_bytes_read.Empty()) { - // TODO(reznor): logger.atVerbose().log("No more data for outgoing payload - // %s, closing InputStream.", this); + NEARBY_LOGS(INFO) << "No more data for outgoing payload " << this + << ", closing InputStream."; input_stream->Close(); return {}; @@ -113,9 +114,6 @@ class OutgoingStreamInternalPayload : public InternalPayload { InputStream* stream = payload_.AsStream(); if (stream) stream->Close(); } - - private: - static constexpr std::int64_t kChunkSize = Pipe::kChunkSize; }; class IncomingStreamInternalPayload : public InternalPayload { @@ -129,7 +127,7 @@ class IncomingStreamInternalPayload : public InternalPayload { std::int64_t GetTotalSize() const override { return -1; } - ByteArray DetachNextChunk() override { return {}; } + ByteArray DetachNextChunk(int chunk_size) override { return {}; } Exception AttachNextChunk(const ByteArray& chunk) override { if (chunk.Empty()) { @@ -158,11 +156,11 @@ class OutgoingFileInternalPayload : public InternalPayload { std::int64_t GetTotalSize() const override { return total_size_; } - ByteArray DetachNextChunk() override { + ByteArray DetachNextChunk(int chunk_size) override { InputFile* file = payload_.AsFile(); if (!file) return {}; - ExceptionOr bytes_read = file->Read(kChunkSize); + ExceptionOr bytes_read = file->Read(chunk_size); if (!bytes_read.ok()) { return {}; } @@ -190,7 +188,6 @@ class OutgoingFileInternalPayload : public InternalPayload { private: std::int64_t total_size_; - static constexpr std::int64_t kChunkSize = 64 * 1024; }; class IncomingFileInternalPayload : public InternalPayload { @@ -207,7 +204,7 @@ class IncomingFileInternalPayload : public InternalPayload { std::int64_t GetTotalSize() const override { return total_size_; } - ByteArray DetachNextChunk() override { return {}; } + ByteArray DetachNextChunk(int chunk_size) override { return {}; } Exception AttachNextChunk(const ByteArray& chunk) override { if (chunk.Empty()) { diff --git a/cpp/core/internal/mediums/webrtc.cc b/cpp/core/internal/mediums/webrtc.cc index ead57cfc..cfa7034d 100644 --- a/cpp/core/internal/mediums/webrtc.cc +++ b/cpp/core/internal/mediums/webrtc.cc @@ -26,6 +26,7 @@ #include "platform/public/logging.h" #include "platform/public/mutex_lock.h" #include "location/nearby/mediums/proto/web_rtc_signaling_frames.pb.h" +#include "absl/container/flat_hash_map.h" #include "absl/strings/str_cat.h" #include "absl/time/time.h" #include "webrtc/api/jsep.h" @@ -53,7 +54,23 @@ WebRtc::~WebRtc() { restart_receive_messages_executor_.Shutdown(); single_thread_executor_.Shutdown(); - Disconnect(); + // Disconnect will also erase the connection info from map. Use a separate + // set to save the connection ids to avoid the iterator violation issue. + absl::flat_hash_set connection_ids; + for (auto& item : accepting_map_) { + connection_ids.emplace(item.first); + } + for (const auto& connection_id : connection_ids) { + Disconnect(Role::kOfferer, connection_id); + } + connection_ids.clear(); + for (auto& item : connecting_map_) { + connection_ids.emplace(item.first); + } + for (const auto& connection_id : connection_ids) { + Disconnect(Role::kAnswerer, connection_id); + } + connection_ids.clear(); } const std::string WebRtc::GetDefaultCountryCode() { @@ -64,8 +81,9 @@ bool WebRtc::IsAvailable() { return medium_.IsValid(); } bool WebRtc::IsAcceptingConnections(const std::string& service_id) { MutexLock lock(&mutex_); - // TODO(hais): refractor the implementation with maps. - return role_ == Role::kOfferer; + ConnectionInfo* connection_info = + GetConnectionInfo(Role::kOfferer, service_id); + return connection_info && connection_info->self_id.IsValid(); } bool WebRtc::StartAcceptingConnections(const std::string& service_id, @@ -73,10 +91,9 @@ bool WebRtc::StartAcceptingConnections(const std::string& service_id, const LocationHint& location_hint, AcceptedConnectionCallback callback) { if (!IsAvailable()) { - { - MutexLock lock(&mutex_); - LogAndDisconnect("WebRTC is not available for data transfer."); - } + MutexLock lock(&mutex_); + LogAndDisconnect(Role::kOfferer, service_id, + "WebRTC is not available for data transfer."); return false; } @@ -84,35 +101,36 @@ bool WebRtc::StartAcceptingConnections(const std::string& service_id, NEARBY_LOG(WARNING, "Already accepting WebRTC connections."); return false; } - { MutexLock lock(&mutex_); - if (role_ != Role::kNone) { - NEARBY_LOG(WARNING, - "Cannot start accepting WebRTC connections, current role %d", - role_); + accepting_map_.emplace(service_id, + ConnectionInfo{.socket = WebRtcSocketWrapper()}); + ConnectionInfo* connection_info = &accepting_map_[service_id]; + if (!InitWebRtcFlow(Role::kOfferer, self_id, location_hint, service_id)) return false; - } - if (!InitWebRtcFlow(Role::kOfferer, self_id, location_hint)) return false; - - restart_receive_messages_alarm_ = CancelableAlarm( + connection_info->restart_receive_messages_alarm = CancelableAlarm( "restart_receiving_messages_webrtc", std::bind(&WebRtc::RestartReceiveMessages, this, location_hint, service_id), kRestartReceiveMessagesDuration, &restart_receive_messages_executor_); - SessionDescriptionWrapper offer = connection_flow_->CreateOffer(); - pending_local_offer_ = webrtc_frames::EncodeOffer(self_id, offer.GetSdp()); - if (!SetLocalSessionDescription(std::move(offer))) { + SessionDescriptionWrapper offer = + connection_info->connection_flow->CreateOffer(); + connection_info->pending_local_offer = + webrtc_frames::EncodeOffer(self_id, offer.GetSdp()); + if (!SetLocalSessionDescription(std::move(offer), Role::kOfferer, + service_id)) { return false; } // There is no timeout set for the future returned since we do not know how // much time it will take for the two devices to discover each other before // the actual transport can begin. - ListenForWebRtcSocketFuture(connection_flow_->GetDataChannel(), - std::move(callback)); + ListenForWebRtcSocketFuture( + Role::kOfferer, service_id, + connection_info->connection_flow->GetDataChannel(), + std::move(callback)); NEARBY_LOG(INFO, "Started listening for WebRtc connections as %s", self_id.GetId().c_str()); } @@ -123,31 +141,38 @@ bool WebRtc::StartAcceptingConnections(const std::string& service_id, WebRtcSocketWrapper WebRtc::Connect(const PeerId& peer_id, const LocationHint& location_hint) { if (!IsAvailable()) { - Disconnect(); + Disconnect(Role::kAnswerer, peer_id.GetId()); return WebRtcSocketWrapper(); } { MutexLock lock(&mutex_); - if (role_ != Role::kNone) { + if (connecting_map_.contains(peer_id.GetId())) { NEARBY_LOG( - WARNING, - "Cannot connect with WebRtc because we are already acting as %d", - role_); + ERROR, + "Cannot connect with WebRtc because we are already connecting."); return WebRtcSocketWrapper(); } - - peer_id_ = peer_id; - if (!InitWebRtcFlow(Role::kAnswerer, PeerId::FromRandom(), location_hint)) { + connecting_map_.emplace(peer_id.GetId(), + ConnectionInfo{.socket = WebRtcSocketWrapper()}); + ConnectionInfo* connection_info = &connecting_map_[peer_id.GetId()]; + connection_info->peer_id = peer_id; + if (!InitWebRtcFlow(Role::kAnswerer, PeerId::FromRandom(), location_hint, + peer_id.GetId())) { return WebRtcSocketWrapper(); } } - NEARBY_LOG(INFO, "Attempting to make a WebRTC connection to %s.", + NEARBY_LOG(ERROR, "Attempting to make a WebRTC connection to %s.", peer_id.GetId().c_str()); - - Future socket_future = ListenForWebRtcSocketFuture( - connection_flow_->GetDataChannel(), AcceptedConnectionCallback()); + Future socket_future; + { + MutexLock lock(&mutex_); + socket_future = ListenForWebRtcSocketFuture( + Role::kAnswerer, peer_id.GetId(), + connecting_map_[peer_id.GetId()].connection_flow->GetDataChannel(), + AcceptedConnectionCallback()); + } // The two devices have discovered each other, hence we have a timeout for // establishing the transport channel. @@ -158,13 +183,19 @@ WebRtcSocketWrapper WebRtc::Connect(const PeerId& peer_id, socket_future.Get(kDataChannelTimeout); if (result.ok()) return result.result(); - Disconnect(); + Disconnect(Role::kAnswerer, peer_id.GetId()); return WebRtcSocketWrapper(); } -bool WebRtc::SetLocalSessionDescription(SessionDescriptionWrapper sdp) { - if (!connection_flow_->SetLocalSessionDescription(std::move(sdp))) { - LogAndDisconnect("Unable to set local session description"); +bool WebRtc::SetLocalSessionDescription(SessionDescriptionWrapper sdp, + Role role, + const std::string& connection_id) { + ConnectionInfo* connection_info = GetConnectionInfo(role, connection_id); + if (!connection_info) return false; + if (!connection_info->connection_flow->SetLocalSessionDescription( + std::move(sdp))) { + LogAndDisconnect(role, connection_id, + "Unable to set local session description"); return false; } @@ -182,28 +213,35 @@ void WebRtc::StopAcceptingConnections(const std::string& service_id) { { MutexLock lock(&mutex_); - ShutdownSignaling(); + ShutdownSignaling(Role::kOfferer, service_id); } NEARBY_LOG(INFO, "Stopped accepting WebRTC connections"); } Future WebRtc::ListenForWebRtcSocketFuture( + const Role& role, const std::string& connection_id, Future> data_channel_future, AcceptedConnectionCallback callback) { Future socket_future; - auto data_channel_runnable = [this, socket_future, data_channel_future, + auto data_channel_runnable = [this, role, connection_id, socket_future, + data_channel_future, callback{std::move(callback)}]() mutable { // The overall timeout of creating the socket and data channel is controlled // by the caller of this function. ExceptionOr> res = data_channel_future.Get(); if (res.ok()) { - WebRtcSocketWrapper wrapper = CreateWebRtcSocketWrapper(res.result()); + WebRtcSocketWrapper wrapper = + CreateWebRtcSocketWrapper(role, connection_id, res.result()); callback.accepted_cb(wrapper); { MutexLock lock(&mutex_); - socket_ = wrapper; + ConnectionInfo* connection_info = + GetConnectionInfo(role, connection_id); + if (connection_info) { + connection_info->socket = wrapper; + } } socket_future.Set(wrapper); } else { @@ -219,274 +257,370 @@ Future WebRtc::ListenForWebRtcSocketFuture( } WebRtcSocketWrapper WebRtc::CreateWebRtcSocketWrapper( + const Role& role, const std::string& connection_id, rtc::scoped_refptr data_channel) { if (data_channel == nullptr) { return WebRtcSocketWrapper(); } auto socket = std::make_unique("WebRtcSocket", data_channel); - socket->SetOnSocketClosedListener( - {[this]() { OffloadFromSignalingThread([this]() { Disconnect(); }); }}); + socket->SetOnSocketClosedListener({[this, role, connection_id]() { + OffloadFromSignalingThread( + [this, role, connection_id]() { Disconnect(role, connection_id); }); + }}); return WebRtcSocketWrapper(std::move(socket)); } -bool WebRtc::InitWebRtcFlow(Role role, const PeerId& self_id, - const LocationHint& location_hint) { - role_ = role; - self_id_ = self_id; +bool WebRtc::InitWebRtcFlow(const Role& role, const PeerId& self_id, + const LocationHint& location_hint, + const std::string& connection_id) { + ConnectionInfo* connection_info = GetConnectionInfo(role, connection_id); + if (!connection_info) return false; + connection_info->self_id = self_id; - if (connection_flow_) { + if (connection_info->connection_flow) { LogAndShutdownSignaling( + role, connection_id, "Tried to initialize WebRTC without shutting down the previous " "connection"); return false; } - if (signaling_messenger_) { + if (connection_info->signaling_messenger) { LogAndShutdownSignaling( + role, connection_id, "Tried to initialize WebRTC without shutting down signaling messenger"); return false; } + connection_info->signaling_messenger = + medium_.GetSignalingMessenger(self_id.GetId(), location_hint); + auto signaling_message_callback = std::bind( + [this](ByteArray message, Role role, const std::string& connection_id) { + OffloadFromSignalingThread([this, message{std::move(message)}, + role{role}, + connection_id{connection_id}]() { + ProcessSignalingMessage(role, connection_id, message); + }); + }, + std::placeholders::_1, role, connection_id); - signaling_messenger_ = - medium_.GetSignalingMessenger(self_id_.GetId(), location_hint); - auto signaling_message_callback = [this](ByteArray message) { - OffloadFromSignalingThread([this, message{std::move(message)}]() { - ProcessSignalingMessage(message); - }); - }; - - if (!signaling_messenger_->IsValid() || - !signaling_messenger_->StartReceivingMessages( + if (!connection_info->signaling_messenger->IsValid() || + !connection_info->signaling_messenger->StartReceivingMessages( signaling_message_callback)) { - DisconnectLocked(); + LogAndDisconnect(role, connection_id, + "Could not receive from signaling messenger."); return false; } - if (role_ == Role::kAnswerer && - !signaling_messenger_->SendMessage( - peer_id_.GetId(), + if (role == Role::kAnswerer && + !connection_info->signaling_messenger->SendMessage( + connection_info->peer_id.GetId(), webrtc_frames::EncodeReadyForSignalingPoke(self_id))) { - LogAndDisconnect(absl::StrCat("Could not send signaling poke to peer ", - peer_id_.GetId())); + LogAndDisconnect(Role::kAnswerer, connection_info->peer_id.GetId(), + absl::StrCat("Could not send signaling poke to peer ", + connection_info->peer_id.GetId())); return false; } - connection_flow_ = ConnectionFlow::Create(GetLocalIceCandidateListener(), - GetDataChannelListener(), medium_); - if (!connection_flow_) return false; + connection_info->connection_flow = ConnectionFlow::Create( + GetLocalIceCandidateListener(role, connection_id), + GetDataChannelListener(role, connection_id), medium_); + if (!connection_info->connection_flow) { + LogAndDisconnect(role, connection_id, "Failed to create connection flow"); + return false; + } return true; } void WebRtc::OnLocalIceCandidate( + const Role& role, const std::string& connection_id, const webrtc::IceCandidateInterface* local_ice_candidate) { ::location::nearby::mediums::IceCandidate ice_candidate = webrtc_frames::EncodeIceCandidate(*local_ice_candidate); - OffloadFromSignalingThread([this, ice_candidate{std::move(ice_candidate)}]() { + OffloadFromSignalingThread([this, ice_candidate{std::move(ice_candidate)}, + role{role}, connection_id{connection_id}]() { MutexLock lock(&mutex_); - if (IsSignaling()) { - signaling_messenger_->SendMessage( - peer_id_.GetId(), webrtc_frames::EncodeIceCandidates( - self_id_, {std::move(ice_candidate)})); + ConnectionInfo* connection_info = GetConnectionInfo(role, connection_id); + if (IsSignaling(role, connection_id)) { + if (connection_info && connection_info->signaling_messenger) { + connection_info->signaling_messenger->SendMessage( + connection_info->peer_id.GetId(), + webrtc_frames::EncodeIceCandidates(connection_info->self_id, + {std::move(ice_candidate)})); + } else { + connection_info->pending_local_ice_candidates.push_back( + std::move(ice_candidate)); + } } else { - pending_local_ice_candidates_.push_back(std::move(ice_candidate)); + connection_info->pending_local_ice_candidates.push_back( + std::move(ice_candidate)); } }); } -LocalIceCandidateListener WebRtc::GetLocalIceCandidateListener() { - return {std::bind(&WebRtc::OnLocalIceCandidate, this, std::placeholders::_1)}; +LocalIceCandidateListener WebRtc::GetLocalIceCandidateListener( + const Role& role, const std::string& connection_id) { + return {std::bind(&WebRtc::OnLocalIceCandidate, this, role, connection_id, + std::placeholders::_1)}; } -void WebRtc::OnDataChannelClosed() { - OffloadFromSignalingThread([this]() { +void WebRtc::OnDataChannelClosed(const Role& role, + const std::string& connection_id) { + OffloadFromSignalingThread([this, role, connection_id]() { MutexLock lock(&mutex_); - LogAndDisconnect("WebRTC data channel closed"); + LogAndDisconnect(role, connection_id, "WebRTC data channel closed"); }); } -void WebRtc::OnDataChannelMessageReceived(const ByteArray& message) { - OffloadFromSignalingThread([this, message]() { - MutexLock lock(&mutex_); - if (!socket_.IsValid()) { - LogAndDisconnect("Received a data channel message without a socket"); - return; +void WebRtc::OnDataChannelMessageReceived(const Role& role, + const std::string& connection_id, + const ByteArray& message) { + OffloadFromSignalingThread([this, role, connection_id, message]() { + { + MutexLock lock(&mutex_); + ConnectionInfo* connection_info = GetConnectionInfo(role, connection_id); + if (!connection_info) return; + if (!connection_info->socket.IsValid()) { + LogAndDisconnect(role, connection_id, + "Received a data channel message without a socket"); + return; + } + connection_info->socket.NotifyDataChannelMsgReceived(message); } - - socket_.NotifyDataChannelMsgReceived(message); }); } -void WebRtc::OnDataChannelBufferedAmountChanged() { - OffloadFromSignalingThread([this]() { - MutexLock lock(&mutex_); - if (!socket_.IsValid()) { - LogAndDisconnect("Data channel buffer changed without a socket"); - return; +void WebRtc::OnDataChannelBufferedAmountChanged( + const Role& role, const std::string& connection_id) { + OffloadFromSignalingThread([this, role, connection_id]() { + { + MutexLock lock(&mutex_); + ConnectionInfo* connection_info = GetConnectionInfo(role, connection_id); + if (!connection_info) return; + if (!connection_info->socket.IsValid()) { + LogAndDisconnect(role, connection_id, + "Data channel buffer changed without a socket"); + return; + } + connection_info->socket.NotifyDataChannelBufferedAmountChanged(); } - - socket_.NotifyDataChannelBufferedAmountChanged(); }); } -DataChannelListener WebRtc::GetDataChannelListener() { +DataChannelListener WebRtc::GetDataChannelListener( + const Role& role, const std::string& connection_id) { return { - .data_channel_closed_cb = std::bind(&WebRtc::OnDataChannelClosed, this), - .data_channel_message_received_cb = std::bind( - &WebRtc::OnDataChannelMessageReceived, this, std::placeholders::_1), + .data_channel_closed_cb = + std::bind(&WebRtc::OnDataChannelClosed, this, role, connection_id), + .data_channel_message_received_cb = + std::bind(&WebRtc::OnDataChannelMessageReceived, this, role, + connection_id, std::placeholders::_1), .data_channel_buffered_amount_changed_cb = - std::bind(&WebRtc::OnDataChannelBufferedAmountChanged, this), + std::bind(&WebRtc::OnDataChannelBufferedAmountChanged, this, role, + connection_id), }; } -bool WebRtc::IsSignaling() { - return (role_ != Role::kNone && self_id_.IsValid() && peer_id_.IsValid()); +bool WebRtc::IsSignaling(const Role& role, const std::string& connection_id) { + ConnectionInfo* connection_info = GetConnectionInfo(role, connection_id); + if (!connection_info) return false; + return (connection_info->self_id.IsValid() && + connection_info->peer_id.IsValid()); } -void WebRtc::ProcessSignalingMessage(const ByteArray& message) { +void WebRtc::ProcessSignalingMessage(const Role& role, + const std::string& connection_id, + const ByteArray& message) { MutexLock lock(&mutex_); + ConnectionInfo* connection_info = GetConnectionInfo(role, connection_id); + if (!connection_info) return; - if (!connection_flow_) { - LogAndDisconnect("Received WebRTC frame before signaling was started"); + if (!connection_info->connection_flow) { + LogAndDisconnect(role, connection_id, + "Received WebRTC frame before signaling was started"); return; } location::nearby::mediums::WebRtcSignalingFrame frame; if (!frame.ParseFromString(std::string(message))) { - LogAndDisconnect("Failed to parse signaling message"); + LogAndDisconnect(role, connection_id, "Failed to parse signaling message"); return; } if (!frame.has_sender_id()) { - LogAndDisconnect("Invalid WebRTC frame: Sender ID is missing"); + LogAndDisconnect(role, connection_id, + "Invalid WebRTC frame: Sender ID is missing"); return; } - if (frame.has_ready_for_signaling_poke() && !peer_id_.IsValid()) { - peer_id_ = PeerId(frame.sender_id().id()); + if (frame.has_ready_for_signaling_poke() && + !connection_info->peer_id.IsValid()) { + connection_info->peer_id = PeerId(frame.sender_id().id()); NEARBY_LOG(INFO, "Peer %s is ready for signaling", - peer_id_.GetId().c_str()); + connection_info->peer_id.GetId().c_str()); } - if (!IsSignaling()) { + if (!IsSignaling(role, connection_id)) { NEARBY_LOG(INFO, "Ignoring WebRTC frame: we are not currently listening for " "signaling messages"); return; } - if (frame.sender_id().id() != peer_id_.GetId()) { + if (frame.sender_id().id() != connection_info->peer_id.GetId()) { NEARBY_LOG( INFO, "Ignoring WebRTC frame: we are only listening for another peer."); return; } if (frame.has_ready_for_signaling_poke()) { - SendOfferAndIceCandidatesToPeer(); + SendOfferAndIceCandidatesToPeer(connection_id); } else if (frame.has_offer()) { - connection_flow_->OnOfferReceived( + DCHECK(role == Role::kAnswerer); + connection_info->connection_flow->OnOfferReceived( SessionDescriptionWrapper(webrtc_frames::DecodeOffer(frame).release())); - SendAnswerToPeer(); + SendAnswerToPeer(connection_id); } else if (frame.has_answer()) { - connection_flow_->OnAnswerReceived(SessionDescriptionWrapper( - webrtc_frames::DecodeAnswer(frame).release())); + DCHECK(role == Role::kOfferer); + connection_info->connection_flow->OnAnswerReceived( + SessionDescriptionWrapper( + webrtc_frames::DecodeAnswer(frame).release())); } else if (frame.has_ice_candidates()) { - if (!connection_flow_->OnRemoteIceCandidatesReceived( + if (!connection_info->connection_flow->OnRemoteIceCandidatesReceived( webrtc_frames::DecodeIceCandidates(frame))) { - LogAndDisconnect("Could not add remote ice candidates."); + LogAndDisconnect(role, connection_id, + "Could not add remote ice candidates."); } } } -void WebRtc::SendOfferAndIceCandidatesToPeer() { - if (pending_local_offer_.Empty()) { +void WebRtc::SendOfferAndIceCandidatesToPeer(const std::string& service_id) { + ConnectionInfo* connection_info = + GetConnectionInfo(Role::kOfferer, service_id); + if (!connection_info) return; + if (connection_info->pending_local_offer.Empty()) { LogAndDisconnect( + Role::kOfferer, service_id, "Unable to send pending offer to remote peer: local offer not set"); return; } - if (!signaling_messenger_->SendMessage(peer_id_.GetId(), - pending_local_offer_)) { - LogAndDisconnect("Failed to send local offer via signaling messenger"); + if (!connection_info->signaling_messenger->SendMessage( + connection_info->peer_id.GetId(), + connection_info->pending_local_offer)) { + LogAndDisconnect(Role::kOfferer, service_id, + "Failed to send local offer via signaling messenger"); return; } - pending_local_offer_ = ByteArray(); + connection_info->pending_local_offer = ByteArray(); - if (!pending_local_ice_candidates_.empty()) { - signaling_messenger_->SendMessage( - peer_id_.GetId(), + if (!connection_info->pending_local_ice_candidates.empty()) { + connection_info->signaling_messenger->SendMessage( + connection_info->peer_id.GetId(), webrtc_frames::EncodeIceCandidates( - self_id_, std::move(pending_local_ice_candidates_))); + connection_info->self_id, + std::move(connection_info->pending_local_ice_candidates))); } } -void WebRtc::SendAnswerToPeer() { - SessionDescriptionWrapper answer = connection_flow_->CreateAnswer(); +void WebRtc::SendAnswerToPeer(const std::string& peer_id) { + ConnectionInfo* connection_info = GetConnectionInfo(Role::kAnswerer, peer_id); + if (!connection_info) return; + SessionDescriptionWrapper answer = + connection_info->connection_flow->CreateAnswer(); ByteArray answer_message( - webrtc_frames::EncodeAnswer(self_id_, answer.GetSdp())); + webrtc_frames::EncodeAnswer(connection_info->self_id, answer.GetSdp())); - if (!SetLocalSessionDescription(std::move(answer))) return; + if (!SetLocalSessionDescription(std::move(answer), Role::kAnswerer, peer_id)) + return; - if (!signaling_messenger_->SendMessage(peer_id_.GetId(), answer_message)) { - LogAndDisconnect("Failed to send local answer via signaling messenger"); + if (!connection_info->signaling_messenger->SendMessage( + connection_info->peer_id.GetId(), answer_message)) { + LogAndDisconnect(Role::kAnswerer, peer_id, + "Failed to send local answer via signaling messenger"); return; } } -void WebRtc::LogAndDisconnect(const std::string& error_message) { - NEARBY_LOG(WARNING, "Disconnecting WebRTC : %s", error_message.c_str()); - DisconnectLocked(); +void WebRtc::LogAndDisconnect(const Role& role, + const std::string& connection_id, + const std::string& error_message) { + NEARBY_LOG(WARNING, + "Disconnecting WebRTC role: %d, connection id: %s, msg: %s", role, + connection_id.c_str(), error_message.c_str()); + DisconnectLocked(role, connection_id); } -void WebRtc::LogAndShutdownSignaling(const std::string& error_message) { - NEARBY_LOG(WARNING, "Stopping WebRTC signaling : %s", error_message.c_str()); - ShutdownSignaling(); +void WebRtc::LogAndShutdownSignaling(const Role& role, + const std::string& connection_id, + const std::string& error_message) { + NEARBY_LOG(WARNING, "Stopping WebRTC role: %d, connection id: %s, msg: %s", + role, connection_id.c_str(), error_message.c_str()); + ShutdownSignaling(role, connection_id); } -void WebRtc::ShutdownSignaling() { - role_ = Role::kNone; - self_id_ = PeerId(); - peer_id_ = PeerId(); - pending_local_offer_ = ByteArray(); - pending_local_ice_candidates_.clear(); - - if (restart_receive_messages_alarm_.IsValid()) { - restart_receive_messages_alarm_.Cancel(); - restart_receive_messages_alarm_ = CancelableAlarm(); +void WebRtc::ShutdownSignaling(const Role& role, + const std::string& connection_id) { + ConnectionInfo* connection_info = GetConnectionInfo(role, connection_id); + if (!connection_info) { + return; } - if (signaling_messenger_) { - signaling_messenger_->StopReceivingMessages(); - signaling_messenger_.reset(); + connection_info->self_id = PeerId(); + connection_info->peer_id = PeerId(); + connection_info->pending_local_offer = ByteArray(); + connection_info->pending_local_ice_candidates.clear(); + + if (connection_info->restart_receive_messages_alarm.IsValid()) { + connection_info->restart_receive_messages_alarm.Cancel(); + connection_info->restart_receive_messages_alarm = CancelableAlarm(); } - if (!socket_.IsValid()) ShutdownIceCandidateCollection(); + if (connection_info->signaling_messenger) { + connection_info->signaling_messenger->StopReceivingMessages(); + connection_info->signaling_messenger.reset(); + } + + if (!connection_info->socket.IsValid()) + ShutdownIceCandidateCollection(role, connection_id); } -void WebRtc::Disconnect() { +void WebRtc::Disconnect(const Role& role, const std::string& connection_id) { MutexLock lock(&mutex_); - DisconnectLocked(); + DisconnectLocked(role, connection_id); } -void WebRtc::DisconnectLocked() { - ShutdownSignaling(); - ShutdownWebRtcSocket(); - ShutdownIceCandidateCollection(); -} +void WebRtc::DisconnectLocked(const Role& role, + const std::string& connection_id) { + ShutdownSignaling(role, connection_id); + ShutdownWebRtcSocket(role, connection_id); + ShutdownIceCandidateCollection(role, connection_id); -void WebRtc::ShutdownWebRtcSocket() { - if (socket_.IsValid()) { - socket_.Close(); - socket_ = WebRtcSocketWrapper(); + if (role == Role::kOfferer && accepting_map_.contains(connection_id)) { + accepting_map_.erase(connection_id); + } else if (role == Role::kAnswerer && + connecting_map_.contains(connection_id)) { + connecting_map_.erase(connection_id); } } -void WebRtc::ShutdownIceCandidateCollection() { - if (connection_flow_) { - connection_flow_->Close(); - connection_flow_.reset(); +void WebRtc::ShutdownWebRtcSocket(const Role& role, + const std::string& connection_id) { + ConnectionInfo* connection_info = GetConnectionInfo(role, connection_id); + if (connection_info && connection_info->socket.IsValid()) { + connection_info->socket.Close(); + connection_info->socket = WebRtcSocketWrapper(); + } +} + +void WebRtc::ShutdownIceCandidateCollection(const Role& role, + const std::string& connection_id) { + ConnectionInfo* connection_info = GetConnectionInfo(role, connection_id); + if (connection_info && connection_info->connection_flow) { + connection_info->connection_flow->Close(); + connection_info->connection_flow.reset(); } } @@ -505,25 +639,49 @@ void WebRtc::RestartReceiveMessages(const LocationHint& location_hint, NEARBY_LOG(INFO, "Restarting listening for receiving signaling messages."); { MutexLock lock(&mutex_); - signaling_messenger_->StopReceivingMessages(); + ConnectionInfo* connection_info = + GetConnectionInfo(Role::kOfferer, service_id); + if (!connection_info) { + NEARBY_LOG(ERROR, + "Can't find connection info in RestartReceiveMessages for %s", + service_id.c_str()); + return; + } + connection_info->signaling_messenger->StopReceivingMessages(); - signaling_messenger_ = - medium_.GetSignalingMessenger(self_id_.GetId(), location_hint); + connection_info->signaling_messenger = medium_.GetSignalingMessenger( + connection_info->self_id.GetId(), location_hint); - auto signaling_message_callback = [this](ByteArray message) { - OffloadFromSignalingThread([this, message{std::move(message)}]() { - ProcessSignalingMessage(message); - }); - }; + auto signaling_message_callback = std::bind( + [this](ByteArray message, const Role& role, + const std::string& connection_id) { + OffloadFromSignalingThread([this, message{std::move(message)}, + role{role}, + connection_id{connection_id}]() { + ProcessSignalingMessage(role, connection_id, message); + }); + }, + std::placeholders::_1, Role::kOfferer, service_id); - if (!signaling_messenger_->IsValid() || - !signaling_messenger_->StartReceivingMessages( + if (!connection_info->signaling_messenger->IsValid() || + !connection_info->signaling_messenger->StartReceivingMessages( signaling_message_callback)) { - DisconnectLocked(); + DisconnectLocked(Role::kOfferer, service_id); } } } +WebRtc::ConnectionInfo* WebRtc::GetConnectionInfo( + const Role& role, const std::string& connection_id) { + if (role == Role::kOfferer && accepting_map_.contains(connection_id)) { + return &accepting_map_[connection_id]; + } else if (role == Role::kAnswerer && + connecting_map_.contains(connection_id)) { + return &connecting_map_[connection_id]; + } + return nullptr; +} + } // namespace mediums } // namespace connections } // namespace nearby diff --git a/cpp/core/internal/mediums/webrtc.h b/cpp/core/internal/mediums/webrtc.h index 1cc1867e..4252efcb 100644 --- a/cpp/core/internal/mediums/webrtc.h +++ b/cpp/core/internal/mediums/webrtc.h @@ -37,6 +37,7 @@ #include "platform/public/single_thread_executor.h" #include "platform/public/webrtc.h" #include "location/nearby/mediums/proto/web_rtc_signaling_frames.pb.h" +#include "absl/container/flat_hash_map.h" #include "absl/container/flat_hash_set.h" #include "webrtc/api/data_channel_interface.h" #include "webrtc/api/jsep.h" @@ -99,65 +100,102 @@ class WebRtc { kAnswerer = 2, }; - bool InitWebRtcFlow(Role role, const PeerId& self_id, - const LocationHint& location_hint) + struct ConnectionInfo { + std::unique_ptr connection_flow; + std::unique_ptr signaling_messenger; + WebRtcSocketWrapper socket; + CancelableAlarm restart_receive_messages_alarm; + + PeerId self_id; + PeerId peer_id; + ByteArray pending_local_offer; + std::vector<::location::nearby::mediums::IceCandidate> + pending_local_ice_candidates; + }; + + bool InitWebRtcFlow(const Role& role, const PeerId& self_id, + const LocationHint& location_hint, + const std::string& connection_id) ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); Future ListenForWebRtcSocketFuture( + const Role& role, const std::string& connection_id, Future> data_channel_future, AcceptedConnectionCallback callback); WebRtcSocketWrapper CreateWebRtcSocketWrapper( + const Role& role, const std::string& connection_id, rtc::scoped_refptr data_channel); - LocalIceCandidateListener GetLocalIceCandidateListener(); + LocalIceCandidateListener GetLocalIceCandidateListener( + const Role& role, const std::string& connection_id); void OnLocalIceCandidate( + const Role& role, const std::string& connection_id, const webrtc::IceCandidateInterface* local_ice_candidate); - DataChannelListener GetDataChannelListener(); - void OnDataChannelClosed(); - void OnDataChannelMessageReceived(const ByteArray& message); - void OnDataChannelBufferedAmountChanged(); + DataChannelListener GetDataChannelListener(const Role& role, + const std::string& connection_id); + void OnDataChannelClosed(const Role& role, const std::string& connection_id); + void OnDataChannelMessageReceived(const Role& role, + const std::string& connection_id, + const ByteArray& message); + void OnDataChannelBufferedAmountChanged(const Role& role, + const std::string& connection_id); // Runs on @MainThread and |single_thread_executor_|. - bool SetLocalSessionDescription(SessionDescriptionWrapper sdp) + bool SetLocalSessionDescription(SessionDescriptionWrapper sdp, Role role, + const std::string& connection_id) ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); // Runs on |single_thread_executor_|. - bool IsSignaling() ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); + bool IsSignaling(const Role& role, const std::string& connection_id) + ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); // Runs on |single_thread_executor_|. - void ProcessSignalingMessage(const ByteArray& message) + void ProcessSignalingMessage(const Role& role, + const std::string& connection_id, + const ByteArray& message) ABSL_LOCKS_EXCLUDED(mutex_); // Runs on |single_thread_executor_|. - void SendOfferAndIceCandidatesToPeer() ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); + void SendOfferAndIceCandidatesToPeer(const std::string& service_id) + ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); // Runs on |single_thread_executor_|. - void SendAnswerToPeer() ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); + void SendAnswerToPeer(const std::string& peer_id) + ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); // Runs on @MainThread and |single_thread_executor_|. - void LogAndDisconnect(const std::string& error_message) + void LogAndDisconnect(const Role& role, const std::string& connection_id, + const std::string& error_message) ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); // Runs on @MainThread. - void Disconnect() ABSL_LOCKS_EXCLUDED(mutex_); + void Disconnect(const Role& role, const std::string& connection_id) + ABSL_LOCKS_EXCLUDED(mutex_); // Runs on @MainThread and |single_thread_executor_|. - void DisconnectLocked() ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); + void DisconnectLocked(const Role& role, const std::string& connection_id) + ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); - void LogAndShutdownSignaling(const std::string& error_message) + void LogAndShutdownSignaling(const Role& role, + const std::string& connection_id, + const std::string& error_message) ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); // Runs on @MainThread and |single_thread_executor_|. - void ShutdownSignaling() ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); + void ShutdownSignaling(const Role& role, const std::string& connection_id) + ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); // Runs on @MainThread and |single_thread_executor_|. - void ShutdownWebRtcSocket() ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); + void ShutdownWebRtcSocket(const Role& role, const std::string& connection_id) + ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); // Runs on @MainThread and |single_thread_executor_|. - void ShutdownIceCandidateCollection(); + void ShutdownIceCandidateCollection(const Role& role, + const std::string& connection_id) + ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); void OffloadFromSignalingThread(Runnable runnable); @@ -166,26 +204,27 @@ class WebRtc { const std::string& service_id) ABSL_LOCKS_EXCLUDED(mutex_); + void PrintStatus(const std::string& func); + + ConnectionInfo* GetConnectionInfo(const Role& role, + const std::string& connection_id) + ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); + Mutex mutex_; - Role role_ ABSL_GUARDED_BY(mutex_) = Role::kNone; - PeerId self_id_ ABSL_GUARDED_BY(mutex_); - PeerId peer_id_ ABSL_GUARDED_BY(mutex_); - ByteArray pending_local_offer_ ABSL_GUARDED_BY(mutex_); - std::vector<::location::nearby::mediums::IceCandidate> - pending_local_ice_candidates_ ABSL_GUARDED_BY(mutex_); - WebRtcMedium medium_; - std::unique_ptr connection_flow_; - std::unique_ptr signaling_messenger_ - ABSL_GUARDED_BY(mutex_); - WebRtcSocketWrapper socket_ ABSL_GUARDED_BY(mutex_); SingleThreadExecutor single_thread_executor_; // Restarts the signaling messenger for receiving messages. ScheduledExecutor restart_receive_messages_executor_; - CancelableAlarm restart_receive_messages_alarm_; + + // Use service_id as key for accepting connections. + absl::flat_hash_map accepting_map_ + ABSL_GUARDED_BY(mutex_); + // Use remote peer_id as key for connecting connections. + absl::flat_hash_map connecting_map_ + ABSL_GUARDED_BY(mutex_); }; } // namespace mediums diff --git a/cpp/core/internal/mediums/webrtc_test.cc b/cpp/core/internal/mediums/webrtc_test.cc index 96f40868..0572c7f4 100644 --- a/cpp/core/internal/mediums/webrtc_test.cc +++ b/cpp/core/internal/mediums/webrtc_test.cc @@ -62,7 +62,8 @@ TEST_F(WebRtcTest, StartAcceptingConnectionTwice) { EXPECT_FALSE(webrtc.StartAcceptingConnections( service_id, self_id, location_hint, {mock_accepted_callback_.AsStdFunction()})); - EXPECT_TRUE(webrtc.IsAcceptingConnections(std::string{})); + EXPECT_TRUE(webrtc.IsAcceptingConnections(service_id)); + EXPECT_FALSE(webrtc.IsAcceptingConnections(std::string{})); } // Tests the flow when the device tries to connect but the data channel times @@ -99,7 +100,7 @@ TEST_F(WebRtcTest, StartAcceptingConnection_ThenConnect) { {mock_accepted_callback_.AsStdFunction()})); WebRtcSocketWrapper wrapper = webrtc.Connect(PeerId("random_peer_id"), location_hint); - EXPECT_TRUE(webrtc.IsAcceptingConnections(std::string{})); + EXPECT_TRUE(webrtc.IsAcceptingConnections(service_id)); EXPECT_FALSE(wrapper.IsValid()); EXPECT_FALSE(webrtc.StartAcceptingConnections( service_id, self_id, location_hint, @@ -122,8 +123,9 @@ TEST_F(WebRtcTest, StartAndStopAcceptingConnections) { ASSERT_TRUE(webrtc.StartAcceptingConnections( service_id, self_id, location_hint, {mock_accepted_callback_.AsStdFunction()})); + EXPECT_TRUE(webrtc.IsAcceptingConnections(service_id)); webrtc.StopAcceptingConnections(service_id); - EXPECT_FALSE(webrtc.IsAcceptingConnections(std::string{})); + EXPECT_FALSE(webrtc.IsAcceptingConnections(service_id)); } // Tests the flow when the device tries to connect to two different peers @@ -144,11 +146,8 @@ TEST_F(WebRtcTest, ConnectTwice) { connected.Set(receiver_socket.IsValid()); }}); - using MockAcceptedCallback = - testing::MockFunction; - testing::StrictMock mock_accepted_callback_; device_c.StartAcceptingConnections(service_id, other_id, location_hint, - {mock_accepted_callback_.AsStdFunction()}); + {[](WebRtcSocketWrapper wrapper) {}}); sender_socket = sender.Connect(self_id, location_hint); EXPECT_TRUE(sender_socket.IsValid()); @@ -157,8 +156,11 @@ TEST_F(WebRtcTest, ConnectTwice) { ASSERT_TRUE(devices_connected.ok()); EXPECT_TRUE(devices_connected.result()); - WebRtcSocketWrapper socket = sender.Connect(other_id, location_hint); - EXPECT_FALSE(socket.IsValid()); + WebRtcSocketWrapper socket = + sender.Connect(other_id, location_hint); + EXPECT_TRUE(socket.IsValid()); + socket.Close(); + EXPECT_TRUE(receiver_socket.IsValid()); EXPECT_TRUE(sender_socket.IsValid()); diff --git a/cpp/core/internal/payload_manager.cc b/cpp/core/internal/payload_manager.cc index 38c29a3b..a471c291 100644 --- a/cpp/core/internal/payload_manager.cc +++ b/cpp/core/internal/payload_manager.cc @@ -16,6 +16,7 @@ #include #include +#include #include #include #include @@ -89,8 +90,9 @@ bool PayloadManager::SendPayloadLoop( // This will block if there is no data to transfer. // It will resume when new data arrives, or if Close() is called. + int chunk_size = GetOptimalChunkSize(available_endpoint_ids); ByteArray next_chunk = - pending_payload.GetInternalPayload()->DetachNextChunk(); + pending_payload.GetInternalPayload()->DetachNextChunk(chunk_size); if (shutdown_.Get()) return false; // Save chunk size. We'll need it after we move next_chunk. auto next_chunk_size = next_chunk.size(); @@ -497,6 +499,15 @@ SingleThreadExecutor* PayloadManager::GetOutgoingPayloadExecutor( } } +int PayloadManager::GetOptimalChunkSize(EndpointIds endpoint_ids) { + int minChunkSize = std::numeric_limits::max(); + for (const auto& endpoint_id : endpoint_ids) { + minChunkSize = std::min( + minChunkSize, endpoint_manager_->GetMaxTransmitPacketSize(endpoint_id)); + } + return minChunkSize; +} + PayloadTransferFrame::PayloadHeader PayloadManager::CreatePayloadHeader( const InternalPayload& internal_payload) { PayloadTransferFrame::PayloadHeader payload_header; diff --git a/cpp/core/internal/payload_manager.h b/cpp/core/internal/payload_manager.h index d9bb12bb..fcc1115f 100644 --- a/cpp/core/internal/payload_manager.h +++ b/cpp/core/internal/payload_manager.h @@ -203,6 +203,8 @@ class PayloadManager : public EndpointManager::FrameProcessor { static PayloadProgressInfo::Status PayloadStatusToTransferUpdateStatus( proto::connections::PayloadStatus status); + int GetOptimalChunkSize(EndpointIds endpoint_ids); + PayloadTransferFrame::PayloadHeader CreatePayloadHeader( const InternalPayload& payload); PayloadTransferFrame::PayloadChunk CreatePayloadChunk(std::int64_t offset, diff --git a/cpp/platform/api/BUILD b/cpp/platform/api/BUILD index 313b8f00..e3149177 100644 --- a/cpp/platform/api/BUILD +++ b/cpp/platform/api/BUILD @@ -79,6 +79,7 @@ cc_library( "platform.h", ], visibility = [ + "//googlemac/iPhone/Shared/Nearby/Connections_v2:__subpackages__", "//platform/base:__pkg__", "//platform/impl:__subpackages__", "//platform/public:__pkg__", diff --git a/cpp/platform/impl/g3/webrtc.cc b/cpp/platform/impl/g3/webrtc.cc index 7035d67e..858b23f6 100644 --- a/cpp/platform/impl/g3/webrtc.cc +++ b/cpp/platform/impl/g3/webrtc.cc @@ -61,14 +61,14 @@ void WebRtcMedium::CreatePeerConnection( webrtc::PeerConnectionInterface::RTCConfiguration rtc_config; webrtc::PeerConnectionDependencies dependencies(observer); - signaling_thread_ = rtc::Thread::Create(); - signaling_thread_->SetName("signaling_thread", nullptr); - RTC_CHECK(signaling_thread_->Start()) << "Failed to start thread"; + std::unique_ptr signaling_thread = rtc::Thread::Create(); + signaling_thread->SetName("signaling_thread", nullptr); + RTC_CHECK(signaling_thread->Start()) << "Failed to start thread"; webrtc::PeerConnectionFactoryDependencies factory_dependencies; factory_dependencies.task_queue_factory = webrtc::CreateDefaultTaskQueueFactory(); - factory_dependencies.signaling_thread = signaling_thread_.get(); + factory_dependencies.signaling_thread = signaling_thread.release(); callback(webrtc::CreateModularPeerConnectionFactory( std::move(factory_dependencies)) diff --git a/cpp/platform/impl/g3/webrtc.h b/cpp/platform/impl/g3/webrtc.h index 9cd20215..b362238f 100644 --- a/cpp/platform/impl/g3/webrtc.h +++ b/cpp/platform/impl/g3/webrtc.h @@ -65,7 +65,6 @@ class WebRtcMedium : public api::WebRtcMedium { const connections::LocationHint& location_hint) override; private: - std::unique_ptr signaling_thread_; }; } // namespace g3