From 463d391fcdea65cb799967c8296397180c637e31 Mon Sep 17 00:00:00 2001 From: nohle Date: Fri, 15 Apr 2022 08:14:15 -0700 Subject: [PATCH] [Nearby Connections][C++] Track bandwidth upgrade mediums for each endpoint (1/4) This CL should be a no-op. Provide better access to the service ID during bandwidth upgrade handling. Notably, 1) have the endpoint channel store its associated service ID, and 2) send the service ID on a connect/disconnect. The service ID identifies the application, for example, "NearbySharing". In future CLs, we will use the service IDs to better handle bandwidth-upgrade bookkeeping when multiple services are running simultaneously. Specifically, we will better handle reverting a bandwidth-upgrade medium without interfering with other services. PiperOrigin-RevId: 442019209 --- connections/core.cc | 19 +- connections/implementation/BUILD | 1 + .../implementation/base_endpoint_channel.cc | 14 +- .../implementation/base_endpoint_channel.h | 59 ++---- .../base_endpoint_channel_test.cc | 2 +- .../implementation/base_pcp_handler.cc | 1 + connections/implementation/base_pcp_handler.h | 3 +- .../implementation/base_pcp_handler_test.cc | 2 +- .../implementation/ble_endpoint_channel.cc | 6 +- .../implementation/ble_endpoint_channel.h | 5 +- .../implementation/bluetooth_bwu_handler.cc | 10 +- .../bluetooth_endpoint_channel.cc | 6 +- .../bluetooth_endpoint_channel.h | 3 +- connections/implementation/bwu_manager.cc | 1 + connections/implementation/bwu_manager.h | 1 + .../implementation/encryption_runner_test.cc | 1 + connections/implementation/endpoint_channel.h | 3 + .../endpoint_channel_manager.cc | 1 - .../implementation/endpoint_manager.cc | 35 ++-- connections/implementation/endpoint_manager.h | 8 +- .../implementation/endpoint_manager_test.cc | 5 +- .../mediums/bluetooth_classic.cc | 181 +++++++++--------- .../mediums/bluetooth_classic.h | 25 +-- .../mediums/bluetooth_classic_test.cc | 6 +- connections/implementation/mediums/webrtc.cc | 2 +- connections/implementation/mediums/webrtc.h | 4 +- .../implementation/mediums/webrtc_test.cc | 26 ++- .../implementation/mediums/wifi_lan.cc | 4 +- connections/implementation/mediums/wifi_lan.h | 5 +- .../implementation/mediums/wifi_lan_test.cc | 6 +- .../implementation/p2p_cluster_pcp_handler.cc | 33 ++-- connections/implementation/payload_manager.cc | 1 + connections/implementation/payload_manager.h | 3 +- .../implementation/service_id_constants.h | 42 ++++ .../implementation/webrtc_bwu_handler.cc | 10 +- .../implementation/webrtc_endpoint_channel.cc | 7 +- .../implementation/webrtc_endpoint_channel.h | 5 +- .../implementation/wifi_lan_bwu_handler.cc | 8 +- .../wifi_lan_endpoint_channel.cc | 5 +- .../wifi_lan_endpoint_channel.h | 5 +- 40 files changed, 333 insertions(+), 231 deletions(-) create mode 100644 connections/implementation/service_id_constants.h diff --git a/connections/core.cc b/connections/core.cc index 95e5c889..284d5d27 100644 --- a/connections/core.cc +++ b/connections/core.cc @@ -20,15 +20,27 @@ #include #include "absl/time/clock.h" -#include "internal/platform/feature_flags.h" +#include "connections/implementation/service_id_constants.h" #include "internal/platform/count_down_latch.h" +#include "internal/platform/feature_flags.h" #include "internal/platform/logging.h" namespace location { namespace nearby { namespace connections { + namespace { + constexpr absl::Duration kWaitForDisconnect = absl::Milliseconds(5000); + +// Verify that |service_id| is not empty and will not conflict with any internal +// service ID formats. +void CheckServiceId(absl::string_view service_id) { + assert(!service_id.empty()); + assert(service_id != kUnknownServiceId); + assert(!IsInitiatorUpgradeServiceId(service_id)); +} + } // namespace Core::Core(ServiceControllerRouter* router) : router_(router) {} @@ -52,7 +64,7 @@ void Core::StartAdvertising(absl::string_view service_id, AdvertisingOptions advertising_options, ConnectionRequestInfo info, ResultCallback callback) { - assert(!service_id.empty()); + CheckServiceId(service_id); assert(advertising_options.strategy.IsValid()); router_->StartAdvertising(&client_, service_id, advertising_options, info, @@ -66,7 +78,7 @@ void Core::StopAdvertising(const ResultCallback callback) { void Core::StartDiscovery(absl::string_view service_id, DiscoveryOptions discovery_options, DiscoveryListener listener, ResultCallback callback) { - assert(!service_id.empty()); + CheckServiceId(service_id); assert(discovery_options.strategy.IsValid()); router_->StartDiscovery(&client_, service_id, discovery_options, listener, @@ -76,6 +88,7 @@ void Core::StartDiscovery(absl::string_view service_id, void Core::InjectEndpoint(absl::string_view service_id, OutOfBandConnectionMetadata metadata, ResultCallback callback) { + CheckServiceId(service_id); router_->InjectEndpoint(&client_, service_id, metadata, callback); } diff --git a/connections/implementation/BUILD b/connections/implementation/BUILD index 8b2e4943..8d7b34bf 100644 --- a/connections/implementation/BUILD +++ b/connections/implementation/BUILD @@ -101,6 +101,7 @@ cc_library( "pcp_manager.h", "service_controller.h", "service_controller_router.h", + "service_id_constants.h", "webrtc_bwu_handler.h", "webrtc_endpoint_channel.h", "wifi_lan_bwu_handler.h", diff --git a/connections/implementation/base_endpoint_channel.cc b/connections/implementation/base_endpoint_channel.cc index ea04f9ee..c335e3e9 100644 --- a/connections/implementation/base_endpoint_channel.cc +++ b/connections/implementation/base_endpoint_channel.cc @@ -15,6 +15,7 @@ #include "connections/implementation/base_endpoint_channel.h" #include +#include #include "absl/strings/escaping.h" #include "absl/strings/str_cat.h" @@ -90,11 +91,12 @@ Exception WriteInt(OutputStream* writer, std::int32_t value) { } // namespace -BaseEndpointChannel::BaseEndpointChannel(const std::string& channel_name, +BaseEndpointChannel::BaseEndpointChannel(const std::string& service_id, + const std::string& channel_name, InputStream* reader, OutputStream* writer) : BaseEndpointChannel( - channel_name, reader, writer, + service_id, channel_name, reader, writer, // TODO(edwinwu): Below values should be retrieved from a base socket, // the #MediumSocket in Android counterpart, from which all the // derived medium sockets should dervied, and implement the supported @@ -106,10 +108,12 @@ BaseEndpointChannel::BaseEndpointChannel(const std::string& channel_name, /*try_count*/ 0) {} BaseEndpointChannel::BaseEndpointChannel( - const std::string& channel_name, InputStream* reader, OutputStream* writer, + const std::string& service_id, const std::string& channel_name, + InputStream* reader, OutputStream* writer, proto::connections::ConnectionTechnology technology, proto::connections::ConnectionBand band, int frequency, int try_count) - : channel_name_(channel_name), + : service_id_(service_id), + channel_name_(channel_name), reader_(reader), writer_(writer), technology_(technology), @@ -312,6 +316,8 @@ std::string BaseEndpointChannel::GetType() const { return medium; } +std::string BaseEndpointChannel::GetServiceId() const { return service_id_; } + std::string BaseEndpointChannel::GetName() const { return channel_name_; } int BaseEndpointChannel::GetMaxTransmitPacketSize() const { diff --git a/connections/implementation/base_endpoint_channel.h b/connections/implementation/base_endpoint_channel.h index 563b5d9d..ce613e3d 100644 --- a/connections/implementation/base_endpoint_channel.h +++ b/connections/implementation/base_endpoint_channel.h @@ -37,80 +37,42 @@ namespace connections { class BaseEndpointChannel : public EndpointChannel { public: - BaseEndpointChannel(const std::string& channel_name, InputStream* reader, + BaseEndpointChannel(const std::string& service_id, + const std::string& channel_name, InputStream* reader, OutputStream* writer); - BaseEndpointChannel(const std::string& channel_name, InputStream* reader, + BaseEndpointChannel(const std::string& service_id, + const std::string& channel_name, InputStream* reader, OutputStream* writer, proto::connections::ConnectionTechnology, proto::connections::ConnectionBand band, int frequency, int try_count); ~BaseEndpointChannel() override = default; + // EndpointChannel: ExceptionOr Read() ABSL_LOCKS_EXCLUDED(reader_mutex_, crypto_mutex_, last_read_mutex_) override; - Exception Write(const ByteArray& data) ABSL_LOCKS_EXCLUDED(writer_mutex_, crypto_mutex_) override; - - // Closes this EndpointChannel, without tracking the closure in analytics. void Close() ABSL_LOCKS_EXCLUDED(is_paused_mutex_) override; - - // Closes this EndpointChannel and records the closure with the given reason. void Close(proto::connections::DisconnectionReason reason) override; - - // Returns a one-word type descriptor for the concrete EndpointChannel - // implementation that can be used in log messages; eg: BLUETOOTH, BLE, - // WIFI. std::string GetType() const override; - - // Returns the name of the EndpointChannel. + std::string GetServiceId() const override; std::string GetName() const override; - - // Returns the maximum supported transmit packet size(MTU) for the underlying - // transport. + proto::connections::ConnectionTechnology GetTechnology() const override; + proto::connections::ConnectionBand GetBand() const override; + int GetFrequency() const override; + int GetTryCount() const override; 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. void EnableEncryption(std::shared_ptr context) override; - - // Disables encryption on the EndpointChannel. void DisableEncryption() override; - - // True if the EndpointChannel is currently pausing all writes. bool IsPaused() const ABSL_LOCKS_EXCLUDED(is_paused_mutex_) override; - - // Pauses all writes on this EndpointChannel until resume() is called. void Pause() ABSL_LOCKS_EXCLUDED(is_paused_mutex_) override; - - // Resumes any writes on this EndpointChannel that were suspended when pause() - // was called. void Resume() ABSL_LOCKS_EXCLUDED(is_paused_mutex_) override; - - // Returns the timestamp (returned by ElapsedRealtime) of the last read from - // this endpoint, or -1 if no reads have occurred. absl::Time GetLastReadTimestamp() const ABSL_LOCKS_EXCLUDED(last_read_mutex_) override; - - // Returns the timestamp (returned by ElapsedRealtime) of the last write to - // this endpoint, or -1 if no writes have occurred. absl::Time GetLastWriteTimestamp() const ABSL_LOCKS_EXCLUDED(last_write_mutex_) override; - - // Returns the used technology of this EndpointChannel. - proto::connections::ConnectionTechnology GetTechnology() const override; - - // Returns the used wifi band of this EndpointChannel. - proto::connections::ConnectionBand GetBand() const override; - - // Returns the used wifi frequency of this EndpointChannel. - int GetFrequency() const override; - - // Returns the try count of this EndpointChannel. - int GetTryCount() const override; - void SetAnalyticsRecorder(analytics::AnalyticsRecorder* analytics_recorder, const std::string& endpoint_id) override; @@ -142,6 +104,7 @@ class BaseEndpointChannel : public EndpointChannel { absl::Time last_write_timestamp_ ABSL_GUARDED_BY(last_write_mutex_) = absl::InfinitePast(); + const std::string service_id_; const std::string channel_name_; // The reader and writer are synchronized independently since we can't have diff --git a/connections/implementation/base_endpoint_channel_test.cc b/connections/implementation/base_endpoint_channel_test.cc index 18832994..dab87068 100644 --- a/connections/implementation/base_endpoint_channel_test.cc +++ b/connections/implementation/base_endpoint_channel_test.cc @@ -50,7 +50,7 @@ using EncryptionContext = BaseEndpointChannel::EncryptionContext; class TestEndpointChannel : public BaseEndpointChannel { public: explicit TestEndpointChannel(InputStream* input, OutputStream* output) - : BaseEndpointChannel("channel", input, output) {} + : BaseEndpointChannel("service_id", "channel", input, output) {} MOCK_METHOD(Medium, GetMedium, (), (const override)); MOCK_METHOD(void, CloseImpl, (), (override)); diff --git a/connections/implementation/base_pcp_handler.cc b/connections/implementation/base_pcp_handler.cc index bbb3859d..a2c6ab01 100644 --- a/connections/implementation/base_pcp_handler.cc +++ b/connections/implementation/base_pcp_handler.cc @@ -929,6 +929,7 @@ void BasePcpHandler::OnIncomingFrame(OfflineFrame& frame, } void BasePcpHandler::OnEndpointDisconnect(ClientProxy* client, + const std::string& service_id, const std::string& endpoint_id, CountDownLatch barrier) { if (stop_.Get()) { diff --git a/connections/implementation/base_pcp_handler.h b/connections/implementation/base_pcp_handler.h index f60147d8..a62298e0 100644 --- a/connections/implementation/base_pcp_handler.h +++ b/connections/implementation/base_pcp_handler.h @@ -138,7 +138,8 @@ class BasePcpHandler : public PcpHandler, // Called when an endpoint disconnects while we're waiting for both sides to // approve/reject the connection. // @EndpointManagerThread - void OnEndpointDisconnect(ClientProxy* client, const std::string& endpoint_id, + void OnEndpointDisconnect(ClientProxy* client, const std::string& service_id, + const std::string& endpoint_id, CountDownLatch barrier) override; Pcp GetPcp() const override { return pcp_; } diff --git a/connections/implementation/base_pcp_handler_test.cc b/connections/implementation/base_pcp_handler_test.cc index 8f8e46f0..2fafa616 100644 --- a/connections/implementation/base_pcp_handler_test.cc +++ b/connections/implementation/base_pcp_handler_test.cc @@ -69,7 +69,7 @@ constexpr BooleanMediumSelector kTestCases[] = { class MockEndpointChannel : public BaseEndpointChannel { public: explicit MockEndpointChannel(Pipe* reader, Pipe* writer) - : BaseEndpointChannel("channel", &reader->GetInputStream(), + : BaseEndpointChannel("service_id", "channel", &reader->GetInputStream(), &writer->GetOutputStream()) {} ExceptionOr DoRead() { return BaseEndpointChannel::Read(); } diff --git a/connections/implementation/ble_endpoint_channel.cc b/connections/implementation/ble_endpoint_channel.cc index 58b026c9..7d00301f 100644 --- a/connections/implementation/ble_endpoint_channel.cc +++ b/connections/implementation/ble_endpoint_channel.cc @@ -37,9 +37,11 @@ InputStream* GetInputStreamOrNull(BleSocket& socket) { } // namespace -BleEndpointChannel::BleEndpointChannel(const std::string& channel_name, +BleEndpointChannel::BleEndpointChannel(const std::string& service_id, + const std::string& channel_name, BleSocket socket) - : BaseEndpointChannel(channel_name, GetInputStreamOrNull(socket), + : BaseEndpointChannel(service_id, channel_name, + GetInputStreamOrNull(socket), GetOutputStreamOrNull(socket)), ble_socket_(std::move(socket)) {} diff --git a/connections/implementation/ble_endpoint_channel.h b/connections/implementation/ble_endpoint_channel.h index 5d756a00..6645165a 100644 --- a/connections/implementation/ble_endpoint_channel.h +++ b/connections/implementation/ble_endpoint_channel.h @@ -15,6 +15,8 @@ #ifndef CORE_INTERNAL_BLE_ENDPOINT_CHANNEL_H_ #define CORE_INTERNAL_BLE_ENDPOINT_CHANNEL_H_ +#include + #include "connections/implementation/base_endpoint_channel.h" #include "internal/platform/ble.h" @@ -25,7 +27,8 @@ namespace connections { class BleEndpointChannel final : public BaseEndpointChannel { public: // Creates both outgoing and incoming Ble channels. - BleEndpointChannel(const std::string& channel_name, BleSocket socket); + BleEndpointChannel(const std::string& service_id, + const std::string& channel_name, BleSocket socket); proto::connections::Medium GetMedium() const override; diff --git a/connections/implementation/bluetooth_bwu_handler.cc b/connections/implementation/bluetooth_bwu_handler.cc index 2a8ec0df..ae988adc 100644 --- a/connections/implementation/bluetooth_bwu_handler.cc +++ b/connections/implementation/bluetooth_bwu_handler.cc @@ -47,8 +47,8 @@ void BluetoothBwuHandler::Revert() { void BluetoothBwuHandler::OnIncomingBluetoothConnection( ClientProxy* client, const std::string& service_id, BluetoothSocket socket) { - auto channel = - absl::make_unique(service_id, socket); + auto channel = absl::make_unique( + service_id, /*channel_name=*/service_id, socket); std::unique_ptr connection{ new IncomingSocketConnection{ .socket = @@ -77,7 +77,7 @@ ByteArray BluetoothBwuHandler::InitializeUpgradedMediumForEndpoint( { .accepted_cb = absl::bind_front( &BluetoothBwuHandler::OnIncomingBluetoothConnection, this, - client, service_id), + client), })) { NEARBY_LOGS(ERROR) << "BluetoothBwuHandler couldn't initiate the " "BLUETOOTH upgrade for endpoint " @@ -145,8 +145,8 @@ BluetoothBwuHandler::CreateUpgradedEndpointChannel( << service_name << ", " << mac_address << ") while upgrading endpoint " << endpoint_id; - auto channel = - std::make_unique(service_name, socket); + auto channel = std::make_unique( + service_id, /*channel_name=*/service_name, socket); if (channel == nullptr) { NEARBY_LOGS(ERROR) << "BluetoothBwuHandler failed to create Bluetooth endpoint " diff --git a/connections/implementation/bluetooth_endpoint_channel.cc b/connections/implementation/bluetooth_endpoint_channel.cc index 8d66110d..5b2bfdd4 100644 --- a/connections/implementation/bluetooth_endpoint_channel.cc +++ b/connections/implementation/bluetooth_endpoint_channel.cc @@ -38,8 +38,10 @@ InputStream* GetInputStreamOrNull(BluetoothSocket& socket) { } // namespace BluetoothEndpointChannel::BluetoothEndpointChannel( - const std::string& channel_name, BluetoothSocket socket) - : BaseEndpointChannel(channel_name, GetInputStreamOrNull(socket), + const std::string& service_id, const std::string& channel_name, + BluetoothSocket socket) + : BaseEndpointChannel(service_id, channel_name, + GetInputStreamOrNull(socket), GetOutputStreamOrNull(socket)), bluetooth_socket_(std::move(socket)) {} diff --git a/connections/implementation/bluetooth_endpoint_channel.h b/connections/implementation/bluetooth_endpoint_channel.h index 859f3e4e..d78dee23 100644 --- a/connections/implementation/bluetooth_endpoint_channel.h +++ b/connections/implementation/bluetooth_endpoint_channel.h @@ -27,7 +27,8 @@ namespace connections { class BluetoothEndpointChannel final : public BaseEndpointChannel { public: // Creates both outgoing and incoming BT channels. - BluetoothEndpointChannel(const std::string& channel_name, + BluetoothEndpointChannel(const std::string& service_id, + const std::string& channel_name, BluetoothSocket bluetooth_socket); proto::connections::Medium GetMedium() const override; diff --git a/connections/implementation/bwu_manager.cc b/connections/implementation/bwu_manager.cc index 86ca7c8c..eafc906f 100644 --- a/connections/implementation/bwu_manager.cc +++ b/connections/implementation/bwu_manager.cc @@ -272,6 +272,7 @@ void BwuManager::OnIncomingFrame(OfflineFrame& frame, } void BwuManager::OnEndpointDisconnect(ClientProxy* client, + const std::string& service_id, const std::string& endpoint_id, CountDownLatch barrier) { NEARBY_LOGS(INFO) diff --git a/connections/implementation/bwu_manager.h b/connections/implementation/bwu_manager.h index e3c6df74..c9415736 100644 --- a/connections/implementation/bwu_manager.h +++ b/connections/implementation/bwu_manager.h @@ -91,6 +91,7 @@ class BwuManager : public EndpointManager::FrameProcessor { // Cleans up in-progress upgrades after endpoint disconnection. // @EndpointManagerReaderThread void OnEndpointDisconnect(ClientProxy* client_proxy, + const std::string& service_id, const std::string& endpoint_id, CountDownLatch barrier) override; void Shutdown(); diff --git a/connections/implementation/encryption_runner_test.cc b/connections/implementation/encryption_runner_test.cc index c2440f36..454e0135 100644 --- a/connections/implementation/encryption_runner_test.cc +++ b/connections/implementation/encryption_runner_test.cc @@ -63,6 +63,7 @@ class FakeEndpointChannel : public EndpointChannel { int GetFrequency() const override { return 0; } int GetTryCount() const override { return 0; } std::string GetType() const override { return "fake-channel-type"; } + std::string GetServiceId() const override { return "fake-service-id"; } std::string GetName() const override { return "fake-channel"; } Medium GetMedium() const override { return Medium::BLE; } int GetMaxTransmitPacketSize() const override { return 512; } diff --git a/connections/implementation/endpoint_channel.h b/connections/implementation/endpoint_channel.h index 7d73e7b2..bcfe1fb9 100644 --- a/connections/implementation/endpoint_channel.h +++ b/connections/implementation/endpoint_channel.h @@ -51,6 +51,9 @@ class EndpointChannel { // implementation that can be used in log messages; eg: BLUETOOTH, BLE, WIFI. virtual std::string GetType() const = 0; + // Returns the service that uses this EndpointChannel. + virtual std::string GetServiceId() const = 0; + // Returns the name of the EndpointChannel. virtual std::string GetName() const = 0; diff --git a/connections/implementation/endpoint_channel_manager.cc b/connections/implementation/endpoint_channel_manager.cc index ece10ebc..7d95358f 100644 --- a/connections/implementation/endpoint_channel_manager.cc +++ b/connections/implementation/endpoint_channel_manager.cc @@ -19,7 +19,6 @@ #include #include "absl/time/time.h" -#include "connections/implementation/proto/offline_wire_formats.pb.h" #include "connections/implementation/offline_frames.h" #include "internal/platform/feature_flags.h" #include "internal/platform/logging.h" diff --git a/connections/implementation/endpoint_manager.cc b/connections/implementation/endpoint_manager.cc index 407a3fb0..66be273b 100644 --- a/connections/implementation/endpoint_manager.cc +++ b/connections/implementation/endpoint_manager.cc @@ -14,14 +14,18 @@ #include "connections/implementation/endpoint_manager.h" +#include +#include #include +#include #include -#include "connections/implementation/proto/offline_wire_formats.pb.h" #include "connections/implementation/endpoint_channel.h" #include "connections/implementation/offline_frames.h" -#include "internal/platform/exception.h" +#include "connections/implementation/proto/offline_wire_formats.pb.h" +#include "connections/implementation/service_id_constants.h" #include "internal/platform/count_down_latch.h" +#include "internal/platform/exception.h" #include "internal/platform/logging.h" #include "internal/platform/mutex_lock.h" @@ -29,8 +33,6 @@ namespace location { namespace nearby { namespace connections { -using ::location::nearby::proto::connections::Medium; - constexpr absl::Duration EndpointManager::kProcessEndpointDisconnectionTimeout; constexpr absl::Time EndpointManager::kInvalidTimestamp; @@ -532,6 +534,13 @@ void EndpointManager::RemoveEndpoint(ClientProxy* client, const std::string& endpoint_id, bool notify) { NEARBY_LOGS(INFO) << "RemoveEndpoint for endpoint " << endpoint_id; + + // Grab the service ID before we destroy the channel. + EndpointChannel* channel = + channel_manager_->GetChannelForEndpoint(endpoint_id).get(); + std::string service_id = + channel ? channel->GetServiceId() : std::string(kUnknownServiceId); + // Unregistering from channel_manager_ will also serve to terminate // the dedicated handler and KeepAlive threads we started when we registered // this endpoint. @@ -541,7 +550,7 @@ void EndpointManager::RemoveEndpoint(ClientProxy* client, // up, we can remove the endpoint from ClientProxy after which there // should be no further interactions with the endpoint. // (See b/37352254 for history) - WaitForEndpointDisconnectionProcessing(client, endpoint_id); + WaitForEndpointDisconnectionProcessing(client, service_id, endpoint_id); client->OnDisconnected(endpoint_id, notify); NEARBY_LOGS(INFO) << "Removed endpoint for endpoint " << endpoint_id; @@ -551,11 +560,13 @@ void EndpointManager::RemoveEndpoint(ClientProxy* client, // @EndpointManagerThread void EndpointManager::WaitForEndpointDisconnectionProcessing( - ClientProxy* client, const std::string& endpoint_id) { + ClientProxy* client, const std::string& service_id, + const std::string& endpoint_id) { NEARBY_LOGS(INFO) << "Wait: client=" << client + << "; service_id=" << service_id << "; endpoint_id=" << endpoint_id; - CountDownLatch barrier = - NotifyFrameProcessorsOnEndpointDisconnect(client, endpoint_id); + CountDownLatch barrier = NotifyFrameProcessorsOnEndpointDisconnect( + client, service_id, endpoint_id); NEARBY_LOGS(INFO) << "Waiting for frame processors to disconnect from endpoint " @@ -571,9 +582,11 @@ void EndpointManager::WaitForEndpointDisconnectionProcessing( } CountDownLatch EndpointManager::NotifyFrameProcessorsOnEndpointDisconnect( - ClientProxy* client, const std::string& endpoint_id) { + ClientProxy* client, const std::string& service_id, + const std::string& endpoint_id) { NEARBY_LOGS(INFO) << "NotifyFrameProcessorsOnEndpointDisconnect: client=" - << client << "; endpoint_id=" << endpoint_id; + << client << "; service_id=" << service_id + << "; endpoint_id=" << endpoint_id; MutexLock lock(&frame_processors_lock_); auto total_size = frame_processors_.size(); NEARBY_LOGS(INFO) << "Total frame processors: " << total_size; @@ -586,7 +599,7 @@ CountDownLatch EndpointManager::NotifyFrameProcessorsOnEndpointDisconnect( << "; frame type=" << V1Frame::FrameType_Name(item.first); if (processor) { valid++; - processor->OnEndpointDisconnect(client, endpoint_id, barrier); + processor->OnEndpointDisconnect(client, service_id, endpoint_id, barrier); } else { barrier.CountDown(); } diff --git a/connections/implementation/endpoint_manager.h b/connections/implementation/endpoint_manager.h index 12d47758..6445a46b 100644 --- a/connections/implementation/endpoint_manager.h +++ b/connections/implementation/endpoint_manager.h @@ -16,7 +16,10 @@ #define CORE_INTERNAL_ENDPOINT_MANAGER_H_ #include +#include #include +#include +#include #include "absl/base/thread_annotations.h" #include "absl/container/flat_hash_map.h" @@ -82,6 +85,7 @@ class EndpointManager { // // @EndpointManagerThread virtual void OnEndpointDisconnect(ClientProxy* client, + const std::string& service_id, const std::string& endpoint_id, CountDownLatch barrier) = 0; }; @@ -247,10 +251,12 @@ class EndpointManager { bool notify); void WaitForEndpointDisconnectionProcessing(ClientProxy* client, + const std::string& service_id, const std::string& endpoint_id); CountDownLatch NotifyFrameProcessorsOnEndpointDisconnect( - ClientProxy* client, const std::string& endpoint_id); + ClientProxy* client, const std::string& service_id, + const std::string& endpoint_id); std::vector SendTransferFrameBytes( const std::vector& endpoint_ids, diff --git a/connections/implementation/endpoint_manager_test.cc b/connections/implementation/endpoint_manager_test.cc index 450e59c9..a60ca416 100644 --- a/connections/implementation/endpoint_manager_test.cc +++ b/connections/implementation/endpoint_manager_test.cc @@ -61,6 +61,7 @@ class MockEndpointChannel : public EndpointChannel { MOCK_METHOD(int, GetFrequency, (), (const override)); MOCK_METHOD(int, GetTryCount, (), (const override)); MOCK_METHOD(std::string, GetType, (), (const override)); + MOCK_METHOD(std::string, GetServiceId, (), (const override)); MOCK_METHOD(std::string, GetName, (), (const override)); MOCK_METHOD(Medium, GetMedium, (), (const override)); MOCK_METHOD(int, GetMaxTransmitPacketSize, (), (const override)); @@ -98,8 +99,8 @@ class MockFrameProcessor : public EndpointManager::FrameProcessor { (override)); MOCK_METHOD(void, OnEndpointDisconnect, - (ClientProxy * client, const std::string& endpoint_id, - CountDownLatch barrier), + (ClientProxy * client, const std::string& service_id, + const std::string& endpoint_id, CountDownLatch barrier), (override)); }; diff --git a/connections/implementation/mediums/bluetooth_classic.cc b/connections/implementation/mediums/bluetooth_classic.cc index 9d3eb084..3f4c051f 100644 --- a/connections/implementation/mediums/bluetooth_classic.cc +++ b/connections/implementation/mediums/bluetooth_classic.cc @@ -26,6 +26,21 @@ namespace location { namespace nearby { namespace connections { +namespace { +std::string ScanModeToString(BluetoothAdapter::ScanMode mode) { + switch (mode) { + case BluetoothAdapter::ScanMode::kUnknown: + return "Unknown"; + case BluetoothAdapter::ScanMode::kNone: + return "None"; + case BluetoothAdapter::ScanMode::kConnectable: + return "Connectable"; + case BluetoothAdapter::ScanMode::kConnectableDiscoverable: + return "ConnectableDiscoverable"; + } +} +} // namespace + BluetoothClassic::BluetoothClassic(BluetoothRadio& radio) : radio_(radio) {} BluetoothClassic::~BluetoothClassic() { @@ -56,42 +71,40 @@ bool BluetoothClassic::TurnOnDiscoverability(const std::string& device_name) { MutexLock lock(&mutex_); if (device_name.empty()) { - NEARBY_LOG(INFO, - "Refusing to turn on BT discoverability. Empty device name."); + NEARBY_LOGS(INFO) + << "Refusing to turn on BT discoverability. Empty device name."; return false; } if (!radio_.IsEnabled()) { - NEARBY_LOG(INFO, "Can't turn on BT discoverability. BT is off."); + NEARBY_LOGS(INFO) << "Can't turn on BT discoverability. BT is off."; return false; } if (!IsAvailableLocked()) { - NEARBY_LOG(INFO, "Can't turn on BT discoverability. BT is not available."); + NEARBY_LOGS(INFO) + << "Can't turn on BT discoverability. BT is not available."; return false; } if (IsDiscoverable()) { - NEARBY_LOG(INFO, - "Refusing to turn on BT discoverability; new name='%s'; " - "current name='%s'", - device_name.c_str(), adapter_.GetName().c_str()); + NEARBY_LOGS(INFO) << "Refusing to turn on BT discoverability; new name='" + << device_name << "'; current name='" + << adapter_.GetName() << "'"; return false; } if (!ModifyDeviceName(device_name)) { - NEARBY_LOG(INFO, - "Failed to turn on BT discoverability; " - "failed to set name to %s", - device_name.c_str()); + NEARBY_LOGS(INFO) + << "Failed to turn on BT discoverability; failed to set name to " + << device_name; return false; } if (!ModifyScanMode(ScanMode::kConnectableDiscoverable)) { - NEARBY_LOG(INFO, - "Failed to turn on BT discoverability; " - "failed to set scan_mode to %d", - ScanMode::kConnectableDiscoverable); + NEARBY_LOGS(INFO) << "Failed to turn on BT discoverability; failed to set " + "scan_mode to " + << ScanModeToString(ScanMode::kConnectableDiscoverable); // Don't forget to perform this rollback of the partial state changes we've // made til now. @@ -99,8 +112,8 @@ bool BluetoothClassic::TurnOnDiscoverability(const std::string& device_name) { return false; } - NEARBY_LOG(INFO, "Turned on BT discoverability with device_name=%s", - device_name.c_str()); + NEARBY_LOGS(INFO) << "Turned on BT discoverability with device_name=" + << device_name; return true; } @@ -108,14 +121,14 @@ bool BluetoothClassic::TurnOffDiscoverability() { MutexLock lock(&mutex_); if (!IsDiscoverable()) { - NEARBY_LOG(INFO, "Can't turn off BT discoverability; it is already off"); + NEARBY_LOGS(INFO) << "Can't turn off BT discoverability; it is already off"; return false; } RestoreScanMode(); RestoreDeviceName(); - NEARBY_LOG(INFO, "Turned Bluetooth discoverability off"); + NEARBY_LOGS(INFO) << "Turned Bluetooth discoverability off"; return true; } @@ -148,8 +161,8 @@ bool BluetoothClassic::ModifyScanMode(ScanMode scan_mode) { bool BluetoothClassic::RestoreScanMode() { if (original_scan_mode_ == ScanMode::kUnknown || !adapter_.SetScanMode(original_scan_mode_)) { - NEARBY_LOG(INFO, "Failed to restore original Bluetooth scan mode to %d", - original_scan_mode_); + NEARBY_LOGS(INFO) << "Failed to restore original Bluetooth scan mode to " + << ScanModeToString(original_scan_mode_); return false; } @@ -162,8 +175,8 @@ bool BluetoothClassic::RestoreScanMode() { bool BluetoothClassic::RestoreDeviceName() { if (original_device_name_.empty() || !adapter_.SetName(original_device_name_)) { - NEARBY_LOG(INFO, "Failed to restore original Bluetooth device name to %s", - original_device_name_.c_str()); + NEARBY_LOGS(INFO) << "Failed to restore original Bluetooth device name to " + << original_device_name_; return false; } original_device_name_.clear(); @@ -174,24 +187,25 @@ bool BluetoothClassic::StartDiscovery(DiscoveredDeviceCallback callback) { MutexLock lock(&mutex_); if (!radio_.IsEnabled()) { - NEARBY_LOG(INFO, "Can't discover BT devices because BT isn't enabled."); + NEARBY_LOGS(INFO) << "Can't discover BT devices because BT isn't enabled."; return false; } if (!IsAvailableLocked()) { - NEARBY_LOG(INFO, "Can't discover BT devices because BT isn't available."); + NEARBY_LOGS(INFO) + << "Can't discover BT devices because BT isn't available."; return false; } if (IsDiscovering()) { - NEARBY_LOG(INFO, - "Refusing to start discovery of BT devices because another " - "discovery is already in-progress."); + NEARBY_LOGS(INFO) + << "Refusing to start discovery of BT devices because another " + "discovery is already in-progress."; return false; } if (!medium_.StartDiscovery(callback)) { - NEARBY_LOG(INFO, "Failed to start discovery of BT devices."); + NEARBY_LOGS(INFO) << "Failed to start discovery of BT devices."; return false; } @@ -205,13 +219,13 @@ bool BluetoothClassic::StopDiscovery() { MutexLock lock(&mutex_); if (!IsDiscovering()) { - NEARBY_LOG(INFO, - "Can't stop discovery of BT devices because it never started."); + NEARBY_LOGS(INFO) + << "Can't stop discovery of BT devices because it never started."; return false; } if (!medium_.StopDiscovery()) { - NEARBY_LOG(INFO, "Failed to stop discovery of Bluetooth devices."); + NEARBY_LOGS(INFO) << "Failed to stop discovery of Bluetooth devices."; return false; } @@ -222,51 +236,46 @@ bool BluetoothClassic::StopDiscovery() { bool BluetoothClassic::IsDiscovering() const { return scan_info_.valid; } bool BluetoothClassic::StartAcceptingConnections( - const std::string& service_name, AcceptedConnectionCallback callback) { + const std::string& service_id, AcceptedConnectionCallback callback) { MutexLock lock(&mutex_); - if (service_name.empty()) { - NEARBY_LOG( - INFO, - "Refusing to start accepting BT connections; service name is empty."); + if (service_id.empty()) { + NEARBY_LOGS(INFO) + << "Refusing to start accepting BT connections; service ID is empty."; return false; } if (!radio_.IsEnabled()) { - NEARBY_LOG(INFO, - "Can't create BT server socket [service=%s]; BT is disabled.", - service_name.c_str()); + NEARBY_LOGS(INFO) << "Can't create BT server socket [service=" << service_id + << "]; BT is disabled."; return false; } if (!IsAvailableLocked()) { - NEARBY_LOG( - INFO, - "Can't start accepting BT connections [service=%s]; BT not available.", - service_name.c_str()); + NEARBY_LOGS(INFO) << "Can't start accepting BT connections [service=" + << service_id << "]; BT not available."; return false; } - if (IsAcceptingConnectionsLocked(service_name)) { - NEARBY_LOG(INFO, - "Refusing to start accepting BT connections [service=%s]; BT " - "server is already in-progress with the same name.", - service_name.c_str()); + if (IsAcceptingConnectionsLocked(service_id)) { + NEARBY_LOGS(INFO) + << "Refusing to start accepting BT connections [service=" << service_id + << "]; BT server is already in-progress with the same name."; return false; } BluetoothServerSocket socket = medium_.ListenForService( - service_name, GenerateUuidFromString(service_name)); + service_id, GenerateUuidFromString(service_id)); if (!socket.IsValid()) { - NEARBY_LOG(INFO, "Failed to start accepting Bluetooth connections for %s.", - service_name.c_str()); + NEARBY_LOGS(INFO) << "Failed to start accepting Bluetooth connections for " + << service_id; return false; } // Mark the fact that there's an in-progress Bluetooth server accepting // connections. auto owned_socket = - server_sockets_.emplace(service_name, std::move(socket)).first->second; + server_sockets_.emplace(service_id, std::move(socket)).first->second; // Start the accept loop on a dedicated thread - this stays alive and // listening for new incoming connections until StopAcceptingConnections() is @@ -274,7 +283,7 @@ bool BluetoothClassic::StartAcceptingConnections( accept_loops_runner_.Execute( "bt-accept", [callback = std::move(callback), server_socket = std::move(owned_socket), - service_name]() mutable { + service_id]() mutable { while (true) { BluetoothSocket client_socket = server_socket.Accept(); if (!client_socket.IsValid()) { @@ -282,41 +291,38 @@ bool BluetoothClassic::StartAcceptingConnections( break; } - callback.accepted_cb(std::move(client_socket)); + callback.accepted_cb(service_id, std::move(client_socket)); } }); return true; } -bool BluetoothClassic::IsAcceptingConnections(const std::string& service_name) { +bool BluetoothClassic::IsAcceptingConnections(const std::string& service_id) { MutexLock lock(&mutex_); - return IsAcceptingConnectionsLocked(service_name); + return IsAcceptingConnectionsLocked(service_id); } bool BluetoothClassic::IsAcceptingConnectionsLocked( - const std::string& service_name) { - return server_sockets_.find(service_name) != server_sockets_.end(); + const std::string& service_id) { + return server_sockets_.find(service_id) != server_sockets_.end(); } bool BluetoothClassic::StopAcceptingConnections( - const std::string& service_name) { + const std::string& service_id) { MutexLock lock(&mutex_); - if (service_name.empty()) { - NEARBY_LOG(INFO, - "Unable to stop accepting BT connections because the " - "service_name is empty."); + if (service_id.empty()) { + NEARBY_LOGS(INFO) << "Unable to stop accepting BT connections because the " + "service_id is empty."; return false; } - const auto& it = server_sockets_.find(service_name); + const auto& it = server_sockets_.find(service_id); if (it == server_sockets_.end()) { - NEARBY_LOG(INFO, - "Can't stop accepting BT connections for %s because it was " - "never started.", - service_name.c_str()); + NEARBY_LOGS(INFO) << "Can't stop accepting BT connections for " + << service_id << " because it was never started."; return false; } @@ -337,8 +343,7 @@ bool BluetoothClassic::StopAcceptingConnections( // Finally, close the BluetoothServerSocket. if (!listening_socket.Close().Ok()) { - NEARBY_LOG(INFO, "Failed to close BT server socket for %s.", - service_name.c_str()); + NEARBY_LOGS(INFO) << "Failed to close BT server socket for " << service_id; return false; } @@ -346,12 +351,12 @@ bool BluetoothClassic::StopAcceptingConnections( } BluetoothSocket BluetoothClassic::Connect(BluetoothDevice& bluetooth_device, - const std::string& service_name, + const std::string& service_id, CancellationFlag* cancellation_flag) { for (int attempts_count = 0; attempts_count < kConnectAttemptsLimit; attempts_count++) { auto wrapper_result = - AttemptToConnect(bluetooth_device, service_name, cancellation_flag); + AttemptToConnect(bluetooth_device, service_id, cancellation_flag); if (wrapper_result.IsValid()) { return wrapper_result; } @@ -360,31 +365,29 @@ BluetoothSocket BluetoothClassic::Connect(BluetoothDevice& bluetooth_device, } BluetoothSocket BluetoothClassic::AttemptToConnect( - BluetoothDevice& bluetooth_device, const std::string& service_name, + BluetoothDevice& bluetooth_device, const std::string& service_id, CancellationFlag* cancellation_flag) { MutexLock lock(&mutex_); - NEARBY_LOG(INFO, "BluetoothClassic::Connect: device=%p", &bluetooth_device); + NEARBY_LOGS(INFO) << "BluetoothClassic::Connect: service_id=" << service_id + << ", device=" << &bluetooth_device; // Socket to return. To allow for NRVO to work, it has to be a single object. BluetoothSocket socket; - if (service_name.empty()) { - NEARBY_LOG( - INFO, - "Refusing to create client BT socket because service_name is empty."); + if (service_id.empty()) { + NEARBY_LOGS(INFO) + << "Refusing to create client BT socket because service_id is empty."; return socket; } if (!radio_.IsEnabled()) { - NEARBY_LOG(INFO, - "Can't create client BT socket [service=%s]: BT isn't enabled.", - service_name.c_str()); + NEARBY_LOGS(INFO) << "Can't create client BT socket [service=" << service_id + << "]: BT isn't enabled."; return socket; } if (!IsAvailableLocked()) { - NEARBY_LOG( - INFO, "Can't create client BT socket [service=%s]; BT isn't available.", - service_name.c_str()); + NEARBY_LOGS(INFO) << "Can't create client BT socket [service=" << service_id + << "]; BT isn't available."; return socket; } @@ -394,11 +397,11 @@ BluetoothSocket BluetoothClassic::AttemptToConnect( } socket = medium_.ConnectToService(bluetooth_device, - GenerateUuidFromString(service_name), + GenerateUuidFromString(service_id), cancellation_flag); if (!socket.IsValid()) { - NEARBY_LOG(INFO, "Failed to Connect via BT [service=%s]", - service_name.c_str()); + NEARBY_LOGS(INFO) << "Failed to Connect via BT [service=" << service_id + << "]"; } return socket; diff --git a/connections/implementation/mediums/bluetooth_classic.h b/connections/implementation/mediums/bluetooth_classic.h index 40e974ef..19a040c8 100644 --- a/connections/implementation/mediums/bluetooth_classic.h +++ b/connections/implementation/mediums/bluetooth_classic.h @@ -16,6 +16,7 @@ #define CORE_INTERNAL_MEDIUMS_BLUETOOTH_CLASSIC_H_ #include +#include #include #include "absl/container/flat_hash_map.h" @@ -39,8 +40,8 @@ class BluetoothClassic { // Callback that is invoked when a new connection is accepted. struct AcceptedConnectionCallback { - std::function accepted_cb = - DefaultCallback(); + std::function + accepted_cb = DefaultCallback(); }; explicit BluetoothClassic(BluetoothRadio& bluetooth_radio); @@ -76,22 +77,22 @@ class BluetoothClassic { bool StopDiscovery() ABSL_LOCKS_EXCLUDED(mutex_); // Starts a worker thread, creates a BT server socket, associates it with a - // service name; in a worker thread repeatedly calls ServerSocket::Accept(). + // service ID; in a worker thread repeatedly calls ServerSocket::Accept(). // Any connected sockets returned from Accept() are passed to a callback. // Returns true, if server socket was successfully created, false otherwise. // Called by server. - bool StartAcceptingConnections(const std::string& service_name, + bool StartAcceptingConnections(const std::string& service_id, AcceptedConnectionCallback callback) ABSL_LOCKS_EXCLUDED(mutex_); // Returns true, if object is currently running a Accept() loop. - bool IsAcceptingConnections(const std::string& service_name) + bool IsAcceptingConnections(const std::string& service_id) ABSL_LOCKS_EXCLUDED(mutex_); - // Closes server socket corresponding to a service name. This automatically + // Closes server socket corresponding to a service ID. This automatically // terminates Accept() loop, if it were running. // Called by server. - bool StopAcceptingConnections(const std::string& service_name) + bool StopAcceptingConnections(const std::string& service_id) ABSL_LOCKS_EXCLUDED(mutex_); // Returns true if this object owns a valid platform implementation. @@ -112,7 +113,7 @@ class BluetoothClassic { // Returns socket instance. On success, BluetoothSocket.IsValid() return true. // Called by client. BluetoothSocket Connect(BluetoothDevice& bluetooth_device, - const std::string& service_name, + const std::string& service_id, CancellationFlag* cancellation_flag) ABSL_LOCKS_EXCLUDED(mutex_); @@ -138,7 +139,7 @@ class BluetoothClassic { bool IsAvailableLocked() const ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); // Same as IsAcceptingConnections(), but must be called with mutex_ held. - bool IsAcceptingConnectionsLocked(const std::string& service_name) + bool IsAcceptingConnectionsLocked(const std::string& service_id) ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); // Returns true, if discoverability is enabled with TurnOnDiscoverability(). @@ -165,12 +166,12 @@ class BluetoothClassic { bool IsDiscovering() const ABSL_EXCLUSIVE_LOCKS_REQUIRED(mutex_); // Establishes connection to BT service that was might be started on another - // device with StartAcceptingConnections() using the same service_name. + // device with StartAcceptingConnections() using the same service_id. // Blocks until connection is established, or server-side is terminated. // Returns socket instance. On success, BluetoothSocket.IsValid() return true. // Called by client. BluetoothSocket AttemptToConnect(BluetoothDevice& bluetooth_device, - const std::string& service_name, + const std::string& service_id, CancellationFlag* cancellation_flag); mutable Mutex mutex_; @@ -195,7 +196,7 @@ class BluetoothClassic { // StartAcceptingConnections(). MultiThreadExecutor accept_loops_runner_{kMaxConcurrentAcceptLoops}; - // A map of service Name -> ServerSocket. If map is non-empty, we + // A map of service ID -> ServerSocket. If map is non-empty, we // are currently listening for incoming connections. // BluetoothServerSocket instances are used from accept_loops_runner_, // and thus require pointer stability. diff --git a/connections/implementation/mediums/bluetooth_classic_test.cc b/connections/implementation/mediums/bluetooth_classic_test.cc index f1ef2266..56018bb8 100644 --- a/connections/implementation/mediums/bluetooth_classic_test.cc +++ b/connections/implementation/mediums/bluetooth_classic_test.cc @@ -122,7 +122,8 @@ TEST_P(BluetoothClassicTest, CanConnect) { std::string(kServiceName), { .accepted_cb = - [&socket_for_server, &accept_latch](BluetoothSocket socket) { + [&socket_for_server, &accept_latch](const std::string& service_id, + BluetoothSocket socket) { socket_for_server = std::move(socket); accept_latch.CountDown(); }, @@ -176,7 +177,8 @@ TEST_P(BluetoothClassicTest, CanCancelConnect) { std::string(kServiceName), { .accepted_cb = - [&socket_for_server, &accept_latch](BluetoothSocket socket) { + [&socket_for_server, &accept_latch](const std::string& service_id, + BluetoothSocket socket) { socket_for_server = std::move(socket); accept_latch.CountDown(); }, diff --git a/connections/implementation/mediums/webrtc.cc b/connections/implementation/mediums/webrtc.cc index 8175afa4..46857dcc 100644 --- a/connections/implementation/mediums/webrtc.cc +++ b/connections/implementation/mediums/webrtc.cc @@ -659,7 +659,7 @@ void WebRtc::ProcessDataChannelOpen(const std::string& service_id, accepting_connections_info_.find(service_id); if (accepting_connection_entry != accepting_connections_info_.end()) { accepting_connection_entry->second.accepted_connection_callback.accepted_cb( - socket_wrapper); + service_id, socket_wrapper); return; } diff --git a/connections/implementation/mediums/webrtc.h b/connections/implementation/mediums/webrtc.h index 7369128a..58e46c1e 100644 --- a/connections/implementation/mediums/webrtc.h +++ b/connections/implementation/mediums/webrtc.h @@ -50,8 +50,8 @@ namespace mediums { // Callback that is invoked when a new connection is accepted. struct AcceptedConnectionCallback { - std::function accepted_cb = - DefaultCallback(); + std::function + accepted_cb = DefaultCallback(); }; // Entry point for connecting a data channel between two devices via WebRtc. diff --git a/connections/implementation/mediums/webrtc_test.cc b/connections/implementation/mediums/webrtc_test.cc index 373e444e..dc919cee 100644 --- a/connections/implementation/mediums/webrtc_test.cc +++ b/connections/implementation/mediums/webrtc_test.cc @@ -14,6 +14,8 @@ #include "connections/implementation/mediums/webrtc.h" +#include + #include "gmock/gmock.h" #include "protobuf-matchers/protocol-buffer-matchers.h" #include "gtest/gtest.h" @@ -42,8 +44,8 @@ constexpr FeatureFlags kTestCases[] = { class WebRtcTest : public ::testing::TestWithParam { protected: - using MockAcceptedCallback = - testing::MockFunction; + using MockAcceptedCallback = testing::MockFunction; WebRtcTest() { env_.Stop(); } @@ -66,7 +68,8 @@ TEST_P(WebRtcTest, ConnectBothDevices_ShutdownSignaling_SendData) { receiver.StartAcceptingConnections( service_id, self_id, location_hint, - {[&receiver_socket, connected](WebRtcSocketWrapper wrapper) mutable { + {[&receiver_socket, connected](const std::string& service_id, + WebRtcSocketWrapper wrapper) mutable { receiver_socket = wrapper; connected.Set(receiver_socket.IsValid()); }}); @@ -104,7 +107,8 @@ TEST_P(WebRtcTest, CanCancelConnect) { receiver.StartAcceptingConnections( service_id, self_id, location_hint, - {[&receiver_socket, connected](WebRtcSocketWrapper wrapper) mutable { + {[&receiver_socket, connected](const std::string& service_id, + WebRtcSocketWrapper wrapper) mutable { receiver_socket = wrapper; connected.Set(receiver_socket.IsValid()); }}); @@ -245,13 +249,15 @@ TEST_F(WebRtcTest, ConnectTwice) { receiver.StartAcceptingConnections( service_id, self_id, location_hint, - {[&receiver_socket, connected](WebRtcSocketWrapper wrapper) mutable { + {[&receiver_socket, connected](const std::string& service_id, + WebRtcSocketWrapper wrapper) mutable { receiver_socket = wrapper; connected.Set(receiver_socket.IsValid()); }}); - device_c.StartAcceptingConnections(service_id, other_id, location_hint, - {[](WebRtcSocketWrapper wrapper) {}}); + device_c.StartAcceptingConnections( + service_id, other_id, location_hint, + {[](const std::string& service_id, WebRtcSocketWrapper wrapper) {}}); CancellationFlag flag; sender_socket = sender.Connect(service_id, self_id, location_hint, &flag); @@ -293,7 +299,8 @@ TEST_F(WebRtcTest, ConnectBothDevicesAndAbort) { receiver.StartAcceptingConnections( service_id, self_id, location_hint, - {[&receiver_socket, connected](WebRtcSocketWrapper wrapper) mutable { + {[&receiver_socket, connected](const std::string& service_id, + WebRtcSocketWrapper wrapper) mutable { receiver_socket = wrapper; connected.Set(receiver_socket.IsValid()); }}); @@ -324,7 +331,8 @@ TEST_F(WebRtcTest, ConnectBothDevicesAndSendData) { receiver.StartAcceptingConnections( service_id, self_id, location_hint, - {[&receiver_socket, connected](WebRtcSocketWrapper wrapper) mutable { + {[&receiver_socket, connected](const std::string& service_id, + WebRtcSocketWrapper wrapper) mutable { receiver_socket = wrapper; connected.Set(receiver_socket.IsValid()); }}); diff --git a/connections/implementation/mediums/wifi_lan.cc b/connections/implementation/mediums/wifi_lan.cc index 99b715a7..a03805f2 100644 --- a/connections/implementation/mediums/wifi_lan.cc +++ b/connections/implementation/mediums/wifi_lan.cc @@ -257,14 +257,14 @@ bool WifiLan::StartAcceptingConnections(const std::string& service_id, accept_loops_runner_.Execute( "wifi-lan-accept", [callback = std::move(callback), - server_socket = std::move(owned_server_socket)]() mutable { + server_socket = std::move(owned_server_socket), service_id]() mutable { while (true) { WifiLanSocket client_socket = server_socket.Accept(); if (!client_socket.IsValid()) { server_socket.Close(); break; } - callback.accepted_cb(std::move(client_socket)); + callback.accepted_cb(service_id, std::move(client_socket)); } }); diff --git a/connections/implementation/mediums/wifi_lan.h b/connections/implementation/mediums/wifi_lan.h index c7a3988a..44cb464c 100644 --- a/connections/implementation/mediums/wifi_lan.h +++ b/connections/implementation/mediums/wifi_lan.h @@ -16,6 +16,7 @@ #define CORE_INTERNAL_MEDIUMS_WIFI_LAN_H_ #include +#include #include #include "absl/container/flat_hash_map.h" @@ -37,8 +38,8 @@ class WifiLan { // Callback that is invoked when a new connection is accepted. struct AcceptedConnectionCallback { - std::function accepted_cb = - DefaultCallback(); + std::function + accepted_cb = DefaultCallback(); }; WifiLan() = default; diff --git a/connections/implementation/mediums/wifi_lan_test.cc b/connections/implementation/mediums/wifi_lan_test.cc index 966cf505..883bad98 100644 --- a/connections/implementation/mediums/wifi_lan_test.cc +++ b/connections/implementation/mediums/wifi_lan_test.cc @@ -74,7 +74,8 @@ TEST_P(WifiLanTest, CanConnect) { service_id, { .accepted_cb = - [&socket_for_server, &accept_latch](WifiLanSocket socket) { + [&socket_for_server, &accept_latch](const std::string& service_id, + WifiLanSocket socket) { socket_for_server = std::move(socket); accept_latch.CountDown(); }, @@ -130,7 +131,8 @@ TEST_P(WifiLanTest, CanCancelConnect) { service_id, { .accepted_cb = - [&socket_for_server, &accept_latch](WifiLanSocket socket) { + [&socket_for_server, &accept_latch](const std::string& service_id, + WifiLanSocket socket) { socket_for_server = std::move(socket); accept_latch.CountDown(); }, diff --git a/connections/implementation/p2p_cluster_pcp_handler.cc b/connections/implementation/p2p_cluster_pcp_handler.cc index b380139d..81609e5e 100644 --- a/connections/implementation/p2p_cluster_pcp_handler.cc +++ b/connections/implementation/p2p_cluster_pcp_handler.cc @@ -835,6 +835,7 @@ proto::connections::Medium P2pClusterPcpHandler::StartBluetoothAdvertising( if (!bluetooth_radio_.Enable() || !bluetooth_medium_.StartAcceptingConnections( service_id, {.accepted_cb = [this, client, local_endpoint_info]( + const std::string& service_id, BluetoothSocket socket) { if (!socket.IsValid()) { NEARBY_LOGS(WARNING) @@ -845,14 +846,15 @@ proto::connections::Medium P2pClusterPcpHandler::StartBluetoothAdvertising( } RunOnPcpHandlerThread( "p2p-bt-on-incoming-connection", - [this, client, local_endpoint_info, + [this, client, local_endpoint_info, service_id, socket = std::move(socket)]() RUN_ON_PCP_HANDLER_THREAD() mutable { std::string remote_device_name = socket.GetRemoteDevice().GetName(); auto channel = absl::make_unique( - remote_device_name, socket); + service_id, /*channel_name=*/remote_device_name, + socket); ByteArray remote_device_info{remote_device_name}; OnIncomingConnection( @@ -965,7 +967,8 @@ BasePcpHandler::ConnectImplResult P2pClusterPcpHandler::BluetoothConnectImpl( } auto channel = absl::make_unique( - endpoint->endpoint_id, bluetooth_socket); + endpoint->service_id, /*channel_name=*/endpoint->endpoint_id, + bluetooth_socket); NEARBY_LOGS(VERBOSE) << "Client" << client->GetClientId() << " created Bluetooth endpoint channel to endpoint(id=" << endpoint->endpoint_id << ")."; @@ -1013,7 +1016,8 @@ proto::connections::Medium P2pClusterPcpHandler::StartBleAdvertising( std::string remote_peripheral_name = socket.GetRemotePeripheral().GetName(); auto channel = absl::make_unique( - remote_peripheral_name, socket); + service_id, + /*channel_name=*/remote_peripheral_name, socket); ByteArray remote_peripheral_info = socket.GetRemotePeripheral().GetAdvertisementBytes( service_id); @@ -1047,6 +1051,7 @@ proto::connections::Medium P2pClusterPcpHandler::StartBleAdvertising( if (!bluetooth_radio_.Enable() || !bluetooth_medium_.StartAcceptingConnections( service_id, {.accepted_cb = [this, client, local_endpoint_info]( + const std::string& service_id, BluetoothSocket socket) { if (!socket.IsValid()) { NEARBY_LOGS(WARNING) @@ -1058,14 +1063,15 @@ proto::connections::Medium P2pClusterPcpHandler::StartBleAdvertising( } RunOnPcpHandlerThread( "p2p-bt-on-incoming-connection", - [this, client, local_endpoint_info, + [this, client, local_endpoint_info, service_id, socket = std::move(socket)]() RUN_ON_PCP_HANDLER_THREAD() mutable { std::string remote_device_name = socket.GetRemoteDevice().GetName(); auto channel = absl::make_unique( - remote_device_name, socket); + service_id, + /*channel_name=*/remote_device_name, socket); ByteArray remote_device_info{remote_device_name}; OnIncomingConnection( @@ -1196,8 +1202,8 @@ BasePcpHandler::ConnectImplResult P2pClusterPcpHandler::BleConnectImpl( }; } - auto channel = - absl::make_unique(endpoint->endpoint_id, ble_socket); + auto channel = absl::make_unique( + endpoint->service_id, /*channel_name=*/endpoint->endpoint_id, ble_socket); return BasePcpHandler::ConnectImplResult{ .medium = proto::connections::Medium::BLE, @@ -1218,7 +1224,8 @@ proto::connections::Medium P2pClusterPcpHandler::StartWifiLanAdvertising( if (!wifi_lan_medium_.StartAcceptingConnections( service_id, {.accepted_cb = [this, client, local_endpoint_info, - local_endpoint_id](WifiLanSocket socket) { + local_endpoint_id](const std::string& service_id, + WifiLanSocket socket) { if (!socket.IsValid()) { NEARBY_LOGS(WARNING) << "Invalid socket in accept callback(" @@ -1229,11 +1236,13 @@ proto::connections::Medium P2pClusterPcpHandler::StartWifiLanAdvertising( RunOnPcpHandlerThread( "p2p-wifi-on-incoming-connection", [this, client, local_endpoint_id, local_endpoint_info, + service_id, socket = std::move( socket)]() RUN_ON_PCP_HANDLER_THREAD() mutable { std::string remote_service_name = local_endpoint_id; auto channel = absl::make_unique( - remote_service_name, socket); + service_id, /*channel_name=*/remote_service_name, + socket); ByteArray remote_service_name_byte{remote_service_name}; OnIncomingConnection(client, remote_service_name_byte, @@ -1349,8 +1358,8 @@ BasePcpHandler::ConnectImplResult P2pClusterPcpHandler::WifiLanConnectImpl( }; } - auto channel = - absl::make_unique(endpoint->endpoint_id, socket); + auto channel = absl::make_unique( + endpoint->service_id, /*channel_name=*/endpoint->endpoint_id, socket); NEARBY_LOGS(INFO) << "Client " << client->GetClientId() << " created WifiLan endpoint channel to endpoint(id=" << endpoint->endpoint_id << ")."; diff --git a/connections/implementation/payload_manager.cc b/connections/implementation/payload_manager.cc index 076691cf..446303f5 100644 --- a/connections/implementation/payload_manager.cc +++ b/connections/implementation/payload_manager.cc @@ -492,6 +492,7 @@ void PayloadManager::OnIncomingFrame( } void PayloadManager::OnEndpointDisconnect(ClientProxy* client, + const std::string& service_id, const std::string& endpoint_id, CountDownLatch barrier) { if (shutdown_.Get()) { diff --git a/connections/implementation/payload_manager.h b/connections/implementation/payload_manager.h index 73b7f561..71f2624e 100644 --- a/connections/implementation/payload_manager.h +++ b/connections/implementation/payload_manager.h @@ -62,7 +62,8 @@ class PayloadManager : public EndpointManager::FrameProcessor { proto::connections::Medium current_medium) override; // @EndpointManagerThread - void OnEndpointDisconnect(ClientProxy* client, const std::string& endpoint_id, + void OnEndpointDisconnect(ClientProxy* client, const std::string& service_id, + const std::string& endpoint_id, CountDownLatch barrier) override; void DisconnectFromEndpointManager(); diff --git a/connections/implementation/service_id_constants.h b/connections/implementation/service_id_constants.h new file mode 100644 index 00000000..e7460008 --- /dev/null +++ b/connections/implementation/service_id_constants.h @@ -0,0 +1,42 @@ +// Copyright 2022 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#ifndef CORE_INTERNAL_SERVICE_ID_CONSTANTS_H_ +#define CORE_INTERNAL_SERVICE_ID_CONSTANTS_H_ + +#include "absl/strings/match.h" +#include "absl/strings/string_view.h" + +namespace location { +namespace nearby { +namespace connections { + +constexpr absl::string_view kUnknownServiceId = "UNKNOWN_SERVICE"; + +// A suffix appended to service IDs when initiating a bandwidth upgrade to +// distinguish the mediums from those used for advertising/discovery. +constexpr absl::string_view kInitiatorUpgradeServiceIdPostfix = "_UPGRADE"; + +// Returns true if |service_id| not empty and has the initiator's upgrade +// postfix. +inline bool IsInitiatorUpgradeServiceId(absl::string_view service_id) { + return !service_id.empty() && + absl::EndsWith(service_id, kInitiatorUpgradeServiceIdPostfix); +} + +} // namespace connections +} // namespace nearby +} // namespace location + +#endif // CORE_INTERNAL_SERVICE_ID_CONSTANTS_H_ diff --git a/connections/implementation/webrtc_bwu_handler.cc b/connections/implementation/webrtc_bwu_handler.cc index a1d07324..6742f57e 100644 --- a/connections/implementation/webrtc_bwu_handler.cc +++ b/connections/implementation/webrtc_bwu_handler.cc @@ -49,7 +49,8 @@ void WebrtcBwuHandler::OnIncomingWebrtcConnection( ClientProxy* client, const std::string& upgrade_service_id, mediums::WebRtcSocketWrapper socket) { std::string service_id = Utils::UnwrapUpgradeServiceId(upgrade_service_id); - auto channel = std::make_unique(service_id, socket); + auto channel = std::make_unique( + upgrade_service_id, /*channel_name=*/service_id, socket); auto webrtc_socket = std::make_unique(service_id, socket); std::unique_ptr connection( @@ -80,8 +81,8 @@ ByteArray WebrtcBwuHandler::InitializeUpgradedMediumForEndpoint( upgrade_service_id, self_id, location_hint, { .accepted_cb = absl::bind_front( - &WebrtcBwuHandler::OnIncomingWebrtcConnection, this, client, - upgrade_service_id), + &WebrtcBwuHandler::OnIncomingWebrtcConnection, this, + client), })) { NEARBY_LOG(ERROR, "WebRtcBwuHandler couldn't initiate the WEB_RTC upgrade for " @@ -139,7 +140,8 @@ WebrtcBwuHandler::CreateUpgradedEndpointChannel( peer_id.GetId().c_str(), endpoint_id.c_str()); // Create a new WebRtcEndpointChannel. - auto channel = std::make_unique(service_id, socket); + auto channel = std::make_unique( + service_id, /*channel_name=*/service_id, socket); if (channel == nullptr) { socket.Close(); NEARBY_LOG(ERROR, diff --git a/connections/implementation/webrtc_endpoint_channel.cc b/connections/implementation/webrtc_endpoint_channel.cc index e798482c..5df48712 100644 --- a/connections/implementation/webrtc_endpoint_channel.cc +++ b/connections/implementation/webrtc_endpoint_channel.cc @@ -14,13 +14,16 @@ #include "connections/implementation/webrtc_endpoint_channel.h" +#include + namespace location { namespace nearby { namespace connections { WebRtcEndpointChannel::WebRtcEndpointChannel( - const std::string& channel_name, mediums::WebRtcSocketWrapper socket) - : BaseEndpointChannel(channel_name, &socket.GetInputStream(), + const std::string& service_id, const std::string& channel_name, + mediums::WebRtcSocketWrapper socket) + : BaseEndpointChannel(service_id, channel_name, &socket.GetInputStream(), &socket.GetOutputStream()), webrtc_socket_(std::move(socket)) {} diff --git a/connections/implementation/webrtc_endpoint_channel.h b/connections/implementation/webrtc_endpoint_channel.h index 692e1b81..8aa91201 100644 --- a/connections/implementation/webrtc_endpoint_channel.h +++ b/connections/implementation/webrtc_endpoint_channel.h @@ -15,6 +15,8 @@ #ifndef CORE_INTERNAL_WEBRTC_ENDPOINT_CHANNEL_H_ #define CORE_INTERNAL_WEBRTC_ENDPOINT_CHANNEL_H_ +#include + #include "connections/implementation/base_endpoint_channel.h" #ifdef NO_WEBRTC #include "connections/implementation/mediums/webrtc_socket_stub.h" @@ -28,7 +30,8 @@ namespace connections { class WebRtcEndpointChannel final : public BaseEndpointChannel { public: - WebRtcEndpointChannel(const std::string& channel_name, + WebRtcEndpointChannel(const std::string& service_id, + const std::string& channel_name, mediums::WebRtcSocketWrapper webrtc_socket); proto::connections::Medium GetMedium() const override; diff --git a/connections/implementation/wifi_lan_bwu_handler.cc b/connections/implementation/wifi_lan_bwu_handler.cc index 0135acef..30ae6f7c 100644 --- a/connections/implementation/wifi_lan_bwu_handler.cc +++ b/connections/implementation/wifi_lan_bwu_handler.cc @@ -52,7 +52,7 @@ ByteArray WifiLanBwuHandler::InitializeUpgradedMediumForEndpoint( { .accepted_cb = absl::bind_front( &WifiLanBwuHandler::OnIncomingWifiLanConnection, this, - client, service_id), + client), })) { NEARBY_LOGS(ERROR) << "WifiLanBwuHandler couldn't initiate the WifiLan upgrade for " @@ -140,7 +140,8 @@ WifiLanBwuHandler::CreateUpgradedEndpointChannel( << endpoint_id; // Create a new WifiLanEndpointChannel. - auto channel = absl::make_unique(service_id, socket); + auto channel = absl::make_unique( + service_id, /*channel_name=*/service_id, socket); if (channel == nullptr) { NEARBY_LOGS(ERROR) << "WifiLanBwuHandler failed to create WifiLan endpoint " "channel to the WifiLan service (" @@ -156,7 +157,8 @@ WifiLanBwuHandler::CreateUpgradedEndpointChannel( // Accept Connection Callback. void WifiLanBwuHandler::OnIncomingWifiLanConnection( ClientProxy* client, const std::string& service_id, WifiLanSocket socket) { - auto channel = absl::make_unique(service_id, socket); + auto channel = absl::make_unique( + service_id, /*channel_name=*/service_id, socket); std::unique_ptr connection( new IncomingSocketConnection{ .socket = diff --git a/connections/implementation/wifi_lan_endpoint_channel.cc b/connections/implementation/wifi_lan_endpoint_channel.cc index 85b22f70..5390672a 100644 --- a/connections/implementation/wifi_lan_endpoint_channel.cc +++ b/connections/implementation/wifi_lan_endpoint_channel.cc @@ -23,9 +23,10 @@ namespace location { namespace nearby { namespace connections { -WifiLanEndpointChannel::WifiLanEndpointChannel(const std::string& channel_name, +WifiLanEndpointChannel::WifiLanEndpointChannel(const std::string& service_id, + const std::string& channel_name, WifiLanSocket socket) - : BaseEndpointChannel(channel_name, &socket.GetInputStream(), + : BaseEndpointChannel(service_id, channel_name, &socket.GetInputStream(), &socket.GetOutputStream()), socket_(std::move(socket)) {} diff --git a/connections/implementation/wifi_lan_endpoint_channel.h b/connections/implementation/wifi_lan_endpoint_channel.h index 27c847b0..48ac7f53 100644 --- a/connections/implementation/wifi_lan_endpoint_channel.h +++ b/connections/implementation/wifi_lan_endpoint_channel.h @@ -15,6 +15,8 @@ #ifndef CORE_INTERNAL_WIFI_LAN_ENDPOINT_CHANNEL_H_ #define CORE_INTERNAL_WIFI_LAN_ENDPOINT_CHANNEL_H_ +#include + #include "connections/implementation/base_endpoint_channel.h" #include "internal/platform/wifi_lan.h" @@ -25,7 +27,8 @@ namespace connections { class WifiLanEndpointChannel final : public BaseEndpointChannel { public: // Creates both outgoing and incoming WifiLan channels. - WifiLanEndpointChannel(const std::string& channel_name, WifiLanSocket socket); + WifiLanEndpointChannel(const std::string& service_id, + const std::string& channel_name, WifiLanSocket socket); proto::connections::Medium GetMedium() const override;