mirror of
https://github.com/kidfromjupiter/nearby.git
synced 2026-09-14 22:56:12 -04:00
Merge pull request #26 from hai007/cl-345608764
Roll forward to Cl/345608764
This commit is contained in:
+1
-1
@@ -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",
|
||||
|
||||
@@ -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<EncryptionContext> context) {
|
||||
MutexLock crypto_lock(&crypto_mutex_);
|
||||
|
||||
@@ -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_);
|
||||
|
||||
@@ -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()) {
|
||||
|
||||
@@ -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_;
|
||||
|
||||
@@ -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()) {
|
||||
|
||||
@@ -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_;
|
||||
|
||||
@@ -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<EncryptionContext> context) override {}
|
||||
void DisableEncryption() override {}
|
||||
bool IsPaused() const override { return false; }
|
||||
|
||||
@@ -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<EncryptionContext> context) = 0;
|
||||
|
||||
|
||||
@@ -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<EndpointChannel> channel =
|
||||
channel_manager_->GetChannelForEndpoint(endpoint_id);
|
||||
if (channel == nullptr) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
return channel->GetMaxTransmitPacketSize();
|
||||
}
|
||||
|
||||
std::vector<std::string> EndpointManager::SendPayloadChunk(
|
||||
@@ -441,6 +439,18 @@ std::vector<std::string> 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<std::string> EndpointManager::SendControlMessage(
|
||||
const PayloadTransferFrame::PayloadHeader& header,
|
||||
const PayloadTransferFrame::ControlMessage& control,
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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<EncryptionContext> context), (override));
|
||||
MOCK_METHOD(void, DisableEncryption, (), (override));
|
||||
|
||||
@@ -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.
|
||||
|
||||
@@ -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<ByteArray> bytes_read = input_stream->Read(kChunkSize);
|
||||
ExceptionOr<ByteArray> 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<ByteArray> bytes_read = file->Read(kChunkSize);
|
||||
ExceptionOr<ByteArray> 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()) {
|
||||
|
||||
+339
-181
@@ -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<std::string> 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<WebRtcSocketWrapper> socket_future = ListenForWebRtcSocketFuture(
|
||||
connection_flow_->GetDataChannel(), AcceptedConnectionCallback());
|
||||
Future<WebRtcSocketWrapper> 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<WebRtcSocketWrapper> WebRtc::ListenForWebRtcSocketFuture(
|
||||
const Role& role, const std::string& connection_id,
|
||||
Future<rtc::scoped_refptr<webrtc::DataChannelInterface>>
|
||||
data_channel_future,
|
||||
AcceptedConnectionCallback callback) {
|
||||
Future<WebRtcSocketWrapper> 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<rtc::scoped_refptr<webrtc::DataChannelInterface>> 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<WebRtcSocketWrapper> WebRtc::ListenForWebRtcSocketFuture(
|
||||
}
|
||||
|
||||
WebRtcSocketWrapper WebRtc::CreateWebRtcSocketWrapper(
|
||||
const Role& role, const std::string& connection_id,
|
||||
rtc::scoped_refptr<webrtc::DataChannelInterface> data_channel) {
|
||||
if (data_channel == nullptr) {
|
||||
return WebRtcSocketWrapper();
|
||||
}
|
||||
|
||||
auto socket = std::make_unique<WebRtcSocket>("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
|
||||
|
||||
@@ -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<ConnectionFlow> connection_flow;
|
||||
std::unique_ptr<WebRtcSignalingMessenger> 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<WebRtcSocketWrapper> ListenForWebRtcSocketFuture(
|
||||
const Role& role, const std::string& connection_id,
|
||||
Future<rtc::scoped_refptr<webrtc::DataChannelInterface>>
|
||||
data_channel_future,
|
||||
AcceptedConnectionCallback callback);
|
||||
|
||||
WebRtcSocketWrapper CreateWebRtcSocketWrapper(
|
||||
const Role& role, const std::string& connection_id,
|
||||
rtc::scoped_refptr<webrtc::DataChannelInterface> 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<ConnectionFlow> connection_flow_;
|
||||
std::unique_ptr<WebRtcSignalingMessenger> 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<std::string, ConnectionInfo> accepting_map_
|
||||
ABSL_GUARDED_BY(mutex_);
|
||||
// Use remote peer_id as key for connecting connections.
|
||||
absl::flat_hash_map<std::string, ConnectionInfo> connecting_map_
|
||||
ABSL_GUARDED_BY(mutex_);
|
||||
};
|
||||
|
||||
} // namespace mediums
|
||||
|
||||
@@ -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<void(WebRtcSocketWrapper socket)>;
|
||||
testing::StrictMock<MockAcceptedCallback> 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());
|
||||
|
||||
@@ -16,6 +16,7 @@
|
||||
|
||||
#include <algorithm>
|
||||
#include <cinttypes>
|
||||
#include <limits>
|
||||
#include <memory>
|
||||
#include <string>
|
||||
#include <utility>
|
||||
@@ -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<int>::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;
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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__",
|
||||
|
||||
@@ -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<rtc::Thread> 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))
|
||||
|
||||
@@ -65,7 +65,6 @@ class WebRtcMedium : public api::WebRtcMedium {
|
||||
const connections::LocationHint& location_hint) override;
|
||||
|
||||
private:
|
||||
std::unique_ptr<rtc::Thread> signaling_thread_;
|
||||
};
|
||||
|
||||
} // namespace g3
|
||||
|
||||
Reference in New Issue
Block a user