diff --git a/internal/platform/credential_storage_impl.cc b/internal/platform/credential_storage_impl.cc index 77ccdcb7..4b7621c2 100644 --- a/internal/platform/credential_storage_impl.cc +++ b/internal/platform/credential_storage_impl.cc @@ -23,7 +23,6 @@ namespace nearby { using ::nearby::internal::LocalCredential; using ::nearby::internal::SharedCredential; using ::nearby::presence::CredentialSelector; -using ::nearby::presence::GetPrivateCredentialsResultCallback; using ::nearby::presence::GetPublicCredentialsResultCallback; using ::nearby::presence::PublicCredentialType; @@ -38,10 +37,18 @@ void CredentialStorageImpl::SaveCredentials( public_credential_type, std::move(callback)); } -void CredentialStorageImpl::GetPrivateCredentials( +void CredentialStorageImpl::UpdateLocalCredential( + absl::string_view manager_app_id, absl::string_view account_name, + nearby::internal::LocalCredential credential, + SaveCredentialsResultCallback callback) { + return impl_->UpdateLocalCredential( + manager_app_id, account_name, std::move(credential), std::move(callback)); +} + +void CredentialStorageImpl::GetLocalCredentials( const CredentialSelector& credential_selector, - GetPrivateCredentialsResultCallback callback) { - return impl_->GetPrivateCredentials(credential_selector, std::move(callback)); + GetLocalCredentialsResultCallback callback) { + return impl_->GetLocalCredentials(credential_selector, std::move(callback)); } void CredentialStorageImpl::GetPublicCredentials( diff --git a/internal/platform/credential_storage_impl.h b/internal/platform/credential_storage_impl.h index d9176929..75614582 100644 --- a/internal/platform/credential_storage_impl.h +++ b/internal/platform/credential_storage_impl.h @@ -47,10 +47,15 @@ class CredentialStorageImpl : public api::CredentialStorage { PublicCredentialType public_credential_type, SaveCredentialsResultCallback callback) override; + void UpdateLocalCredential(absl::string_view manager_app_id, + absl::string_view account_name, + nearby::internal::LocalCredential credential, + SaveCredentialsResultCallback callback) override; + // Used to fetch private creds when broadcasting. - void GetPrivateCredentials( + void GetLocalCredentials( const CredentialSelector& credential_selector, - GetPrivateCredentialsResultCallback callback) override; + GetLocalCredentialsResultCallback callback) override; // Used to fetch remote public creds when scanning. void GetPublicCredentials( diff --git a/internal/platform/credential_storage_impl_test.cc b/internal/platform/credential_storage_impl_test.cc index f2ef8479..a3a6339c 100644 --- a/internal/platform/credential_storage_impl_test.cc +++ b/internal/platform/credential_storage_impl_test.cc @@ -36,7 +36,7 @@ using ::nearby::internal::IdentityType; using ::nearby::internal::LocalCredential; using ::nearby::internal::SharedCredential; using ::nearby::presence::CredentialSelector; -using ::nearby::presence::GetPrivateCredentialsResultCallback; +using ::nearby::presence::GetLocalCredentialsResultCallback; using ::nearby::presence::GetPublicCredentialsResultCallback; using ::nearby::presence::PublicCredentialType; using ::nearby::presence::SaveCredentialsResultCallback; @@ -49,8 +49,8 @@ constexpr absl::string_view kManagerAppId = "manager app id"; constexpr absl::string_view kAccountName = "test_account"; // `secret_id` is used to create credentials with different content. -LocalCredential CreatePrivateCredential(absl::string_view secret_id, - IdentityType identity_type) { +LocalCredential CreateLocalCredential(absl::string_view secret_id, + IdentityType identity_type) { LocalCredential private_credential; private_credential.set_secret_id(secret_id); private_credential.set_identity_type(identity_type); @@ -67,10 +67,10 @@ SharedCredential CreatePublicCredential(absl::string_view secret_id, std::vector BuildPrivateCreds(absl::string_view secret_id) { std::vector private_credentials = { - CreatePrivateCredential(secret_id, IdentityType::IDENTITY_TYPE_PRIVATE), - CreatePrivateCredential(secret_id, IdentityType::IDENTITY_TYPE_TRUSTED), - CreatePrivateCredential(secret_id, - IdentityType::IDENTITY_TYPE_PROVISIONED)}; + CreateLocalCredential(secret_id, IdentityType::IDENTITY_TYPE_PRIVATE), + CreateLocalCredential(secret_id, IdentityType::IDENTITY_TYPE_TRUSTED), + CreateLocalCredential(secret_id, + IdentityType::IDENTITY_TYPE_PROVISIONED)}; return private_credentials; } @@ -83,7 +83,7 @@ std::vector BuildPublicCreds(absl::string_view secret_id) { return public_credentials; } -absl::StatusOr> GetPrivateCredentials( +absl::StatusOr> GetLocalCredentials( CredentialStorageImpl& credential_storage, IdentityType identity_type, absl::string_view manager_app_id = kManagerAppId, absl::string_view account_name = kAccountName) { @@ -91,9 +91,9 @@ absl::StatusOr> GetPrivateCredentials( .account_name = std::string(account_name), .identity_type = identity_type}; absl::StatusOr> private_credentials; - credential_storage.GetPrivateCredentials( + credential_storage.GetLocalCredentials( selector, - GetPrivateCredentialsResultCallback{ + GetLocalCredentialsResultCallback{ .credentials_fetched_cb = [&](absl::StatusOr> credentials) { private_credentials = std::move(credentials); @@ -146,8 +146,8 @@ absl::Status SaveCredentials(CredentialStorageImpl& credential_storage, PublicCredentialType::kLocalPublicCredential); } -absl::Status SavePrivateCredentials(CredentialStorageImpl& credential_storage, - absl::string_view secret_id) { +absl::Status SaveLocalCredentials(CredentialStorageImpl& credential_storage, + absl::string_view secret_id) { return SaveCredentials(credential_storage, kManagerAppId, kAccountName, BuildPrivateCreds(secret_id), std::vector(), @@ -162,7 +162,7 @@ absl::Status SavePublicCredentials(CredentialStorageImpl& credential_storage, BuildPublicCreds(secret_id), credential_type); } -TEST(CredentialStorageImplTest, SaveAndGetPrivateCredentials) { +TEST(CredentialStorageImplTest, SaveAndGetLocalCredentials) { std::vector default_private_creds = BuildPrivateCreds(kSecretId); std::vector empty_public_creds; @@ -178,21 +178,49 @@ TEST(CredentialStorageImplTest, SaveAndGetPrivateCredentials) { }}); EXPECT_OK(save_status); - auto fetched_private_credentials = GetPrivateCredentials( + auto fetched_private_credentials = GetLocalCredentials( credential_storage, IdentityType::IDENTITY_TYPE_UNSPECIFIED); ASSERT_OK(fetched_private_credentials); EXPECT_THAT(*fetched_private_credentials, UnorderedPointwise(EqualsProto(), default_private_creds)); } -TEST(CredentialStorageImplTest, ReplaceAndGetPrivateCredentials) { +TEST(CredentialStorageImplTest, UpdateLocalCredential) { + std::vector default_private_creds = + BuildPrivateCreds(kSecretId); + std::vector empty_public_creds; + CredentialStorageImpl credential_storage; + absl::Status update_status = absl::UnknownError(""); + credential_storage.SaveCredentials( + kManagerAppId, kAccountName, default_private_creds, empty_public_creds, + PublicCredentialType::kLocalPublicCredential, + SaveCredentialsResultCallback{ + .credentials_saved_cb = [&](absl::Status status) { + ASSERT_OK(status); + }}); + + // Modify a private credential + default_private_creds[0].mutable_consumed_salts()->insert({1234, true}); + credential_storage.UpdateLocalCredential( + kManagerAppId, kAccountName, default_private_creds[0], + {[&](absl::Status status) { update_status = status; }}); + + EXPECT_OK(update_status); + auto fetched_private_credentials = GetLocalCredentials( + credential_storage, IdentityType::IDENTITY_TYPE_UNSPECIFIED); + ASSERT_OK(fetched_private_credentials); + EXPECT_THAT(*fetched_private_credentials, + UnorderedPointwise(EqualsProto(), default_private_creds)); +} + +TEST(CredentialStorageImplTest, ReplaceAndGetLocalCredentials) { constexpr absl::string_view kAnotherSecretId = "another secret id"; CredentialStorageImpl credential_storage; - EXPECT_OK(SavePrivateCredentials(credential_storage, kSecretId)); - EXPECT_OK(SavePrivateCredentials(credential_storage, kAnotherSecretId)); + EXPECT_OK(SaveLocalCredentials(credential_storage, kSecretId)); + EXPECT_OK(SaveLocalCredentials(credential_storage, kAnotherSecretId)); - auto fetched_private_credentials = GetPrivateCredentials( + auto fetched_private_credentials = GetLocalCredentials( credential_storage, IdentityType::IDENTITY_TYPE_UNSPECIFIED); ASSERT_OK(fetched_private_credentials); EXPECT_THAT( @@ -301,7 +329,7 @@ TEST(CredentialStorageImplTest, SavePrivateAndLocalPublicCredentials) { ASSERT_OK(fetched_public_credentials); EXPECT_THAT(*fetched_public_credentials, UnorderedPointwise(EqualsProto(), public_creds)); - auto fetched_private_credentials = GetPrivateCredentials( + auto fetched_private_credentials = GetLocalCredentials( credential_storage, IdentityType::IDENTITY_TYPE_UNSPECIFIED); ASSERT_OK(fetched_private_credentials); EXPECT_THAT(*fetched_private_credentials, @@ -325,38 +353,38 @@ TEST(CredentialStorageImplTest, SaveCredentialsFailsWhenNoCredentials) { EXPECT_THAT(save_status, StatusIs(absl::StatusCode::kInvalidArgument)); } -TEST(CredentialStorageImplTest, GetPrivateCredentialsFailsWhenNoCredentials) { +TEST(CredentialStorageImplTest, GetLocalCredentialsFailsWhenNoCredentials) { CredentialStorageImpl credential_storage; EXPECT_OK(SavePublicCredentials(credential_storage, PublicCredentialType::kRemotePublicCredential, kSecretId)); - EXPECT_THAT(GetPrivateCredentials(credential_storage, - IdentityType::IDENTITY_TYPE_UNSPECIFIED), + EXPECT_THAT(GetLocalCredentials(credential_storage, + IdentityType::IDENTITY_TYPE_UNSPECIFIED), StatusIs(absl::StatusCode::kNotFound)); } TEST(CredentialStorageImplTest, - GetPrivateCredentialsFailsWhenManagerAppIdDoesNotMatch) { + GetLocalCredentialsFailsWhenManagerAppIdDoesNotMatch) { CredentialStorageImpl credential_storage; EXPECT_OK(SaveCredentials(credential_storage, BuildPrivateCreds(kSecretId), BuildPublicCreds(kSecretId))); - EXPECT_THAT(GetPrivateCredentials(credential_storage, - IdentityType::IDENTITY_TYPE_UNSPECIFIED, - "different manager app id", kAccountName), + EXPECT_THAT(GetLocalCredentials(credential_storage, + IdentityType::IDENTITY_TYPE_UNSPECIFIED, + "different manager app id", kAccountName), StatusIs(absl::StatusCode::kNotFound)); } TEST(CredentialStorageImplTest, - GetPrivateCredentialsFailsWhenAccountNameDoesNotMatch) { + GetLocalCredentialsFailsWhenAccountNameDoesNotMatch) { CredentialStorageImpl credential_storage; EXPECT_OK(SaveCredentials(credential_storage, BuildPrivateCreds(kSecretId), BuildPublicCreds(kSecretId))); - EXPECT_THAT(GetPrivateCredentials(credential_storage, - IdentityType::IDENTITY_TYPE_UNSPECIFIED, - kManagerAppId, "different account name"), + EXPECT_THAT(GetLocalCredentials(credential_storage, + IdentityType::IDENTITY_TYPE_UNSPECIFIED, + kManagerAppId, "different account name"), StatusIs(absl::StatusCode::kNotFound)); } @@ -400,46 +428,45 @@ TEST(CredentialStorageImplTest, GetPublicCredentialsFailsWhenNoCredentials) { class IdentityFilterTest : public testing::TestWithParam {}; -TEST_P(IdentityFilterTest, FilterPrivateCredentialsByIdentityType) { +TEST_P(IdentityFilterTest, FilterLocalCredentialsByIdentityType) { IdentityType identity_type = GetParam(); CredentialStorageImpl credential_storage; - EXPECT_OK(SavePrivateCredentials(credential_storage, kSecretId)); + EXPECT_OK(SaveLocalCredentials(credential_storage, kSecretId)); EXPECT_OK(SavePublicCredentials(credential_storage, PublicCredentialType::kLocalPublicCredential, kSecretId)); auto private_credentials = - GetPrivateCredentials(credential_storage, identity_type); + GetLocalCredentials(credential_storage, identity_type); ASSERT_OK(private_credentials); EXPECT_THAT( *private_credentials, UnorderedPointwise(EqualsProto(), - std::vector{CreatePrivateCredential( - kSecretId, identity_type)})); + std::vector{ + CreateLocalCredential(kSecretId, identity_type)})); } -TEST_P(IdentityFilterTest, - FilterPrivateCredentialsFailsWhenNoCredentialsMatch) { +TEST_P(IdentityFilterTest, FilterLocalCredentialsFailsWhenNoCredentialsMatch) { IdentityType identity_type = GetParam(); // Create a credential of a different identity type than the one we query. IdentityType other_type = identity_type == IdentityType::IDENTITY_TYPE_PRIVATE ? IdentityType::IDENTITY_TYPE_TRUSTED : IdentityType::IDENTITY_TYPE_PRIVATE; std::vector private_creds = { - CreatePrivateCredential(kSecretId, other_type)}; + CreateLocalCredential(kSecretId, other_type)}; CredentialStorageImpl credential_storage; EXPECT_OK(SaveCredentials(credential_storage, private_creds, BuildPublicCreds(kSecretId))); - EXPECT_THAT(GetPrivateCredentials(credential_storage, identity_type), + EXPECT_THAT(GetLocalCredentials(credential_storage, identity_type), StatusIs(absl::StatusCode::kNotFound)); } TEST_P(IdentityFilterTest, FilterPublicCredentialsByIdentityType) { IdentityType identity_type = GetParam(); CredentialStorageImpl credential_storage; - EXPECT_OK(SavePrivateCredentials(credential_storage, kSecretId)); + EXPECT_OK(SaveLocalCredentials(credential_storage, kSecretId)); EXPECT_OK(SavePublicCredentials(credential_storage, PublicCredentialType::kLocalPublicCredential, kSecretId)); diff --git a/internal/platform/implementation/credential_callbacks.h b/internal/platform/implementation/credential_callbacks.h index a201c4e1..fde302c7 100644 --- a/internal/platform/implementation/credential_callbacks.h +++ b/internal/platform/implementation/credential_callbacks.h @@ -74,7 +74,7 @@ struct UpdateRemotePublicCredentialsCallback { absl::AnyInvocable credentials_updated_cb; }; -struct GetPrivateCredentialsResultCallback { +struct GetLocalCredentialsResultCallback { absl::AnyInvocable>)> credentials_fetched_cb; diff --git a/internal/platform/implementation/credential_storage.h b/internal/platform/implementation/credential_storage.h index c682b357..5c0af512 100644 --- a/internal/platform/implementation/credential_storage.h +++ b/internal/platform/implementation/credential_storage.h @@ -34,8 +34,8 @@ class CredentialStorage { using SaveCredentialsResultCallback = ::nearby::presence::SaveCredentialsResultCallback; using CredentialSelector = ::nearby::presence::CredentialSelector; - using GetPrivateCredentialsResultCallback = - ::nearby::presence::GetPrivateCredentialsResultCallback; + using GetLocalCredentialsResultCallback = + ::nearby::presence::GetLocalCredentialsResultCallback; using GetPublicCredentialsResultCallback = ::nearby::presence::GetPublicCredentialsResultCallback; @@ -60,13 +60,20 @@ class CredentialStorage { PublicCredentialType public_credential_type, SaveCredentialsResultCallback callback) = 0; + // Updates the `credential` in the storage. LocalCredential has a + // `secret_id` field, which uniquely identifies the credential. + virtual void UpdateLocalCredential( + absl::string_view manager_app_id, absl::string_view account_name, + nearby::internal::LocalCredential credential, + SaveCredentialsResultCallback callback) = 0; + // Fetches private credentials. // // When `credential_selector.identity_type` is not set (unspecified), then // private credentials with any identity type should be returned. - virtual void GetPrivateCredentials( + virtual void GetLocalCredentials( const CredentialSelector& credential_selector, - GetPrivateCredentialsResultCallback callback) = 0; + GetLocalCredentialsResultCallback callback) = 0; // Fetches public credentials. // diff --git a/internal/platform/implementation/g3/credential_storage_impl.cc b/internal/platform/implementation/g3/credential_storage_impl.cc index 1195c563..b2501128 100644 --- a/internal/platform/implementation/g3/credential_storage_impl.cc +++ b/internal/platform/implementation/g3/credential_storage_impl.cc @@ -14,6 +14,7 @@ #include "internal/platform/implementation/g3/credential_storage_impl.h" +#include #include #include #include @@ -70,15 +71,8 @@ void CredentialStorageImpl::SaveCredentials( << account_name << "], manager app ID:[" << manager_app_id << "]"; absl::MutexLock lock(&private_mutex_); - PrivateCredentialKey key = - CreatePrivateCredentialKey(manager_app_id, account_name); - auto private_result = private_credentials_map_.insert( - std::make_pair(key, private_credentials)); - if (!private_result.second) { - NEARBY_LOGS(WARNING) - << "Credentials already saved in map. Overwriting previous creds!"; - private_credentials_map_[key] = private_credentials; - } + SaveLocalCredentialsLocked(manager_app_id, account_name, + private_credentials); } if (public_credentials.empty()) { @@ -102,30 +96,77 @@ void CredentialStorageImpl::SaveCredentials( } std::move(callback.credentials_saved_cb)(absl::OkStatus()); } +void CredentialStorageImpl::SaveLocalCredentialsLocked( + absl::string_view manager_app_id, absl::string_view account_name, + const std::vector& private_credentials) { + LocalCredentialKey key = + CreateLocalCredentialKey(manager_app_id, account_name); + auto private_result = + private_credentials_map_.insert(std::make_pair(key, private_credentials)); + if (!private_result.second) { + NEARBY_LOGS(WARNING) + << "Credentials already saved in map. Overwriting previous creds!"; + private_credentials_map_[key] = private_credentials; + } +} -void CredentialStorageImpl::GetPrivateCredentials( +void CredentialStorageImpl::UpdateLocalCredential( + absl::string_view manager_app_id, absl::string_view account_name, + LocalCredential credential, SaveCredentialsResultCallback callback) { + NEARBY_LOGS(INFO) << "G3 Update Private Credential for for account: [" + << account_name << "], manager app ID:[" << manager_app_id + << "]"; + absl::MutexLock lock(&private_mutex_); + absl::StatusOr> credentials = + GetLocalCredentialsLocked(CredentialSelector{ + .manager_app_id = std::string(manager_app_id), + .account_name = std::string(account_name), + .identity_type = internal::IDENTITY_TYPE_UNSPECIFIED}); + if (!credentials.ok()) { + NEARBY_LOGS(WARNING) << credentials.status(); + credentials = std::vector(); + } + auto it = std::find_if(credentials->begin(), credentials->end(), + [&](const LocalCredential& a) { + return a.secret_id() == credential.secret_id(); + }); + if (it == credentials->end()) { + credentials->push_back(std::move(credential)); + } else { + *it = std::move(credential); + } + SaveLocalCredentialsLocked(manager_app_id, account_name, *credentials); + callback.credentials_saved_cb(absl::OkStatus()); +} + +void CredentialStorageImpl::GetLocalCredentials( const CredentialSelector& credential_selector, - GetPrivateCredentialsResultCallback callback) { + GetLocalCredentialsResultCallback callback) { NEARBY_LOGS(INFO) << "G3 Get Private Credentials for " << credential_selector; absl::MutexLock lock(&private_mutex_); - PrivateCredentialKey key = CreatePrivateCredentialKey( + std::move(callback.credentials_fetched_cb)( + GetLocalCredentialsLocked(credential_selector)); +} + +absl::StatusOr> +CredentialStorageImpl::GetLocalCredentialsLocked( + const CredentialSelector& credential_selector) { + LocalCredentialKey key = CreateLocalCredentialKey( credential_selector.manager_app_id, credential_selector.account_name); if (private_credentials_map_.find(key) == private_credentials_map_.end()) { NEARBY_LOGS(WARNING) << "There are no Private Credentials stored for key:" << std::get<0>(key) << ", " << std::get<1>(key); - std::move(callback.credentials_fetched_cb)(absl::NotFoundError( - absl::StrFormat("No private credentials for %v", credential_selector))); - return; + return absl::NotFoundError( + absl::StrFormat("No private credentials for %v", credential_selector)); } std::vector private_credentials = private_credentials_map_[key]; FilterIdentityType(private_credentials, credential_selector.identity_type); if (private_credentials.empty()) { - std::move(callback.credentials_fetched_cb)(absl::NotFoundError( - absl::StrFormat("No private credentials for %v", credential_selector))); - return; + return absl::NotFoundError( + absl::StrFormat("No private credentials for %v", credential_selector)); } - std::move(callback.credentials_fetched_cb)(private_credentials); + return private_credentials; } void CredentialStorageImpl::GetPublicCredentials( @@ -155,5 +196,6 @@ void CredentialStorageImpl::GetPublicCredentials( } std::move(callback.credentials_fetched_cb)(public_credentials); } + } // namespace g3 } // namespace nearby diff --git a/internal/platform/implementation/g3/credential_storage_impl.h b/internal/platform/implementation/g3/credential_storage_impl.h index 2d14be72..236a2ddf 100644 --- a/internal/platform/implementation/g3/credential_storage_impl.h +++ b/internal/platform/implementation/g3/credential_storage_impl.h @@ -41,7 +41,7 @@ class CredentialStorageImpl : public api::CredentialStorage { using LocalCredential = ::nearby::internal::LocalCredential; using SharedCredential = ::nearby::internal::SharedCredential; using PublicCredentialType = ::nearby::presence::PublicCredentialType; - using PrivateCredentialKey = std::pair; + using LocalCredentialKey = std::pair; using PublicCredentialKey = std::tuple; @@ -56,10 +56,15 @@ class CredentialStorageImpl : public api::CredentialStorage { PublicCredentialType public_credential_type, SaveCredentialsResultCallback callback) override; + void UpdateLocalCredential(absl::string_view manager_app_id, + absl::string_view account_name, + nearby::internal::LocalCredential credential, + SaveCredentialsResultCallback callback) override; + // Used to fetch private creds when broadcasting. - void GetPrivateCredentials( + void GetLocalCredentials( const CredentialSelector& credential_selector, - GetPrivateCredentialsResultCallback callback) override; + GetLocalCredentialsResultCallback callback) override; // Used to fetch remote public creds when scanning. void GetPublicCredentials( @@ -68,7 +73,7 @@ class CredentialStorageImpl : public api::CredentialStorage { GetPublicCredentialsResultCallback callback) override; private: - PrivateCredentialKey CreatePrivateCredentialKey( + LocalCredentialKey CreateLocalCredentialKey( absl::string_view manager_app_id, absl::string_view account_name) { return std::make_tuple(std::string(manager_app_id), std::string(account_name)); @@ -79,7 +84,13 @@ class CredentialStorageImpl : public api::CredentialStorage { return std::make_tuple(std::string(manager_app_id), std::string(account_name), credential_type); } - absl::flat_hash_map> + absl::StatusOr> GetLocalCredentialsLocked( + const CredentialSelector& credential_selector); + void SaveLocalCredentialsLocked( + absl::string_view manager_app_id, absl::string_view account_name, + const std::vector& private_credentials); + + absl::flat_hash_map> private_credentials_map_; absl::flat_hash_map> public_credentials_map_; diff --git a/presence/implementation/advertisement_factory.cc b/presence/implementation/advertisement_factory.cc index 6327360f..bfad3320 100644 --- a/presence/implementation/advertisement_factory.cc +++ b/presence/implementation/advertisement_factory.cc @@ -15,6 +15,7 @@ #include "presence/implementation/advertisement_factory.h" #include +#include #include #include "absl/status/status.h" @@ -96,11 +97,11 @@ std::string SerializeAction(const Action& action) { absl::StatusOr AdvertisementFactory::CreateAdvertisement( const BaseBroadcastRequest& request, - std::vector& credentials) const { + absl::optional credential) const { AdvertisementData advert = {}; if (absl::holds_alternative( request.variant)) { - return CreateBaseNpAdvertisement(request, credentials); + return CreateBaseNpAdvertisement(request, std::move(credential)); } return advert; } @@ -108,7 +109,7 @@ absl::StatusOr AdvertisementFactory::CreateAdvertisement( absl::StatusOr AdvertisementFactory::CreateBaseNpAdvertisement( const BaseBroadcastRequest& request, - std::vector& credentials) const { + absl::optional credential) const { const auto& presence = absl::get(request.variant); std::string payload; @@ -126,7 +127,7 @@ AdvertisementFactory::CreateBaseNpAdvertisement( return absl::InvalidArgumentError( absl::StrFormat("Unsupported salt size %d", request.salt.size())); } - if (credentials.empty()) { + if (!credential) { return absl::FailedPreconditionError("Missing credentials"); } std::string unencrypted; @@ -137,7 +138,7 @@ AdvertisementFactory::CreateBaseNpAdvertisement( return result; } absl::StatusOr encrypted = - EncryptDataElements(credentials, request.salt, unencrypted); + EncryptDataElements(*credential, request.salt, unencrypted); if (!encrypted.ok()) { return encrypted.status(); } @@ -181,9 +182,8 @@ AdvertisementFactory::CreateBaseNpAdvertisement( .content = payload}; } absl::StatusOr AdvertisementFactory::EncryptDataElements( - std::vector& credentials, absl::string_view salt, + const LocalCredential& credential, absl::string_view salt, absl::string_view data_elements) const { - LocalCredential& credential = credentials.front(); if (credential.metadata_encryption_key().size() != kBaseMetadataSize) { return absl::FailedPreconditionError(absl::StrFormat( "Metadata key size %d, expected %d", diff --git a/presence/implementation/advertisement_factory.h b/presence/implementation/advertisement_factory.h index edde9aa0..69321f54 100644 --- a/presence/implementation/advertisement_factory.h +++ b/presence/implementation/advertisement_factory.h @@ -39,20 +39,19 @@ class AdvertisementFactory { // Returns a BLE advertisement for given `request. absl::StatusOr CreateAdvertisement( const BaseBroadcastRequest& request, - std::vector& credentials) const; + absl::optional credential) const; absl::StatusOr CreateAdvertisement( const BaseBroadcastRequest& request) const { - std::vector empty; - return CreateAdvertisement(request, empty); + return CreateAdvertisement(request, absl::optional()); } private: absl::StatusOr CreateBaseNpAdvertisement( const BaseBroadcastRequest& request, - std::vector& credentials) const; + absl::optional credential) const; absl::StatusOr EncryptDataElements( - std::vector& credentials, absl::string_view salt, + const LocalCredential& credential, absl::string_view salt, absl::string_view data_elements) const; }; diff --git a/presence/implementation/advertisement_factory_test.cc b/presence/implementation/advertisement_factory_test.cc index 3db93440..927a1732 100644 --- a/presence/implementation/advertisement_factory_test.cc +++ b/presence/implementation/advertisement_factory_test.cc @@ -41,7 +41,7 @@ using ::testing::Return; using ::testing::status::StatusIs; #if USE_RUST_LDT == 1 -LocalCredential CreatePrivateCredential(IdentityType identity_type) { +LocalCredential CreateLocalCredential(IdentityType identity_type) { // Values copied from LDT tests ByteArray seed({204, 219, 36, 137, 233, 252, 172, 66, 179, 147, 72, 184, 148, 30, 209, 154, 29, 54, 14, 117, 224, 152, @@ -60,8 +60,6 @@ TEST(AdvertisementFactory, CreateAdvertisementFromPrivateIdentity) { std::string account_name = "Test account"; std::string salt = "AB"; constexpr IdentityType kIdentity = IdentityType::IDENTITY_TYPE_PRIVATE; - std::vector credentials = { - CreatePrivateCredential(kIdentity)}; std::vector data_elements; data_elements.emplace_back(DataElement(ActionBit::kActiveUnlockAction)); Action action = ActionFactory::CreateAction(data_elements); @@ -73,7 +71,8 @@ TEST(AdvertisementFactory, CreateAdvertisementFromPrivateIdentity) { .SetAction(action)); absl::StatusOr result = - AdvertisementFactory().CreateAdvertisement(request, credentials); + AdvertisementFactory().CreateAdvertisement( + request, CreateLocalCredential(kIdentity)); ASSERT_OK(result); EXPECT_FALSE(result->is_extended_advertisement); @@ -85,8 +84,6 @@ TEST(AdvertisementFactory, CreateAdvertisementFromTrustedIdentity) { std::string account_name = "Test account"; std::string salt = "AB"; constexpr IdentityType kIdentity = IdentityType::IDENTITY_TYPE_TRUSTED; - std::vector credentials = { - CreatePrivateCredential(kIdentity)}; std::vector data_elements; data_elements.emplace_back(DataElement(ActionBit::kActiveUnlockAction)); data_elements.emplace_back(DataElement(ActionBit::kFitCastAction)); @@ -99,7 +96,8 @@ TEST(AdvertisementFactory, CreateAdvertisementFromTrustedIdentity) { .SetAction(action)); absl::StatusOr result = - AdvertisementFactory().CreateAdvertisement(request, credentials); + AdvertisementFactory().CreateAdvertisement( + request, CreateLocalCredential(kIdentity)); ASSERT_OK(result); EXPECT_FALSE(result->is_extended_advertisement); @@ -111,8 +109,6 @@ TEST(AdvertisementFactory, CreateAdvertisementFromProvisionedIdentity) { std::string account_name = "Test account"; std::string salt = "AB"; constexpr IdentityType kIdentity = IdentityType::IDENTITY_TYPE_PROVISIONED; - std::vector credentials = { - CreatePrivateCredential(kIdentity)}; std::vector data_elements; data_elements.emplace_back(DataElement(ActionBit::kActiveUnlockAction)); data_elements.emplace_back(DataElement(ActionBit::kFitCastAction)); @@ -125,7 +121,8 @@ TEST(AdvertisementFactory, CreateAdvertisementFromProvisionedIdentity) { .SetAction(action)); absl::StatusOr result = - AdvertisementFactory().CreateAdvertisement(request, credentials); + AdvertisementFactory().CreateAdvertisement( + request, CreateLocalCredential(kIdentity)); ASSERT_OK(result); EXPECT_FALSE(result->is_extended_advertisement); diff --git a/presence/implementation/broadcast_manager.cc b/presence/implementation/broadcast_manager.cc index 3f88b1e4..33e0edcc 100644 --- a/presence/implementation/broadcast_manager.cc +++ b/presence/implementation/broadcast_manager.cc @@ -14,11 +14,18 @@ #include "presence/implementation/broadcast_manager.h" +#include +#include #include +#include +#include #include #include +#include "absl/time/time.h" #include "internal/crypto/random.h" +#include "internal/platform/implementation/system_clock.h" +#include "internal/platform/logging.h" #include "presence/implementation/advertisement_factory.h" namespace nearby { @@ -28,6 +35,42 @@ namespace { using AdvertisingCallback = ::nearby::api::ble_v2::BleMedium::AdvertisingCallback; using AdvertisingSession = ::nearby::api::ble_v2::BleMedium::AdvertisingSession; +using LocalCredential = internal::LocalCredential; + +uint16_t SaltToInt(absl::string_view salt) { + if (salt.length() < 2) return 0; + uint16_t b0 = salt[0]; + uint16_t b1 = salt[1]; + return b0 << 8 | b1; +} +std::string SaltFromInt(uint16_t x) { + std::string salt; + salt.resize(2); + salt[0] = x >> 8 & 0xFF; + salt[1] = x & 0xFF; + return salt; +} + +// Selects a salt that has not been used yet. The salt is added to +// `credential.consumed_salts`. +// We may fail to find an unused salt. In this unlikely event, an already +// consumed salt is returned. +std::string SelectSalt(LocalCredential& credential, + absl::string_view preferred_salt) { + // NP certificate guidelines say that we should try to get an unused salt 128 + // times. + constexpr int kMaxSaltSelectRetries = 128; + + uint16_t s = SaltToInt(preferred_salt); + for (int i = 0; i < kMaxSaltSelectRetries; i++) { + if (!credential.consumed_salts().contains(s)) { + break; + } + s = crypto::RandData(); + } + credential.mutable_consumed_salts()->insert({s, true}); + return SaltFromInt(s); +} } // namespace @@ -63,11 +106,12 @@ void BroadcastManager::FetchCredentials( Advertise(id, broadcast_request, /*credentials=*/{}); return; } - credential_manager_->GetPrivateCredentials( + credential_manager_->GetLocalCredentials( *credential_selector, - GetPrivateCredentialsResultCallback{ + GetLocalCredentialsResultCallback{ .credentials_fetched_cb = - [this, id, broadcast_request = std::move(broadcast_request)]( + [this, id, broadcast_request = std::move(broadcast_request), + selector = *credential_selector]( absl::StatusOr< std::vector<::nearby::internal::LocalCredential>> credentials) { @@ -81,29 +125,67 @@ void BroadcastManager::FetchCredentials( RunOnServiceControllerThread( "advertise-non-public", [this, id, broadcast_request = std::move(broadcast_request), - credentials = std::move(*credentials)]() - ABSL_EXCLUSIVE_LOCKS_REQUIRED(executor_) { - Advertise(id, broadcast_request, credentials); + credentials = std::move(*credentials), + selector = std::move(selector)]() + ABSL_EXCLUSIVE_LOCKS_REQUIRED(executor_) mutable { + absl::optional credential = + Advertise(id, broadcast_request, credentials); + if (credential) { + credential_manager_->UpdateLocalCredential( + selector, std::move(*credential), + {[](absl::Status status) { + if (!status.ok()) { + NEARBY_LOGS(WARNING) + << "Failed to update private " + "credential, status: " + << status; + } + }}); + } }); }}); } -void BroadcastManager::Advertise(BroadcastSessionId id, - BaseBroadcastRequest broadcast_request, - std::vector credentials) { +absl::optional BroadcastManager::SelectCredential( + BaseBroadcastRequest& broadcast_request, + std::vector credentials) { + if (credentials.empty()) { + return absl::optional(); + } + auto credential = + std::min_element(credentials.begin(), credentials.end(), + [](const LocalCredential& a, const LocalCredential& b) { + return a.start_time_millis() < b.start_time_millis(); + }); + if (credential == credentials.end()) { + NEARBY_LOGS(WARNING) << "No active credentials"; + return absl::optional(); + } + std::string salt = SelectSalt(*credential, broadcast_request.salt); + if (salt != broadcast_request.salt) { + NEARBY_LOGS(VERBOSE) << "Changed salt"; + broadcast_request.salt = salt; + } + return *credential; +} + +absl::optional BroadcastManager::Advertise( + BroadcastSessionId id, BaseBroadcastRequest broadcast_request, + std::vector credentials) { auto it = sessions_.find(id); if (it == sessions_.end()) { NEARBY_LOGS(INFO) << "Broadcast session terminated, id: " << id; - return; + return absl::optional(); } + absl::optional credential = + SelectCredential(broadcast_request, std::move(credentials)); absl::StatusOr advertisement = - AdvertisementFactory().CreateAdvertisement(broadcast_request, - credentials); + AdvertisementFactory().CreateAdvertisement(broadcast_request, credential); if (!advertisement.ok()) { NEARBY_LOGS(WARNING) << "Can't create advertisement, reason: " << advertisement.status(); NotifyStartCallbackStatus(id, advertisement.status()); - return; + return absl::optional(); } std::unique_ptr session = mediums_->GetBle().StartAdvertising( @@ -115,9 +197,10 @@ void BroadcastManager::Advertise(BroadcastSessionId id, if (!session) { NotifyStartCallbackStatus(id, absl::InternalError("Can't start advertising")); - return; + return absl::optional(); } it->second.SetAdvertisingSession(std::move(session)); + return credential; } void BroadcastManager::NotifyStartCallbackStatus(BroadcastSessionId id, diff --git a/presence/implementation/broadcast_manager.h b/presence/implementation/broadcast_manager.h index 81eecf62..d9a7fd52 100644 --- a/presence/implementation/broadcast_manager.h +++ b/presence/implementation/broadcast_manager.h @@ -22,6 +22,7 @@ #include "absl/base/thread_annotations.h" #include "absl/strings/string_view.h" +#include "absl/types/optional.h" #include "internal/platform/single_thread_executor.h" #include "internal/proto/credential.pb.h" #include "presence/broadcast_request.h" @@ -83,9 +84,16 @@ class BroadcastManager { void FetchCredentials(BroadcastSessionId id, BaseBroadcastRequest broadcast_request) ABSL_EXCLUSIVE_LOCKS_REQUIRED(*executor_); + absl::optional SelectCredential( + BaseBroadcastRequest& broadcast_request, + std::vector credentials); - void Advertise(BroadcastSessionId id, BaseBroadcastRequest broadcast_request, - std::vector credentials) + // Returns the private credential, if any, selected to generate the + // advertisement. A salt used in the advertisement is added to the returned + // private credential. The caller must save it in the storage. + absl::optional Advertise( + BroadcastSessionId id, BaseBroadcastRequest broadcast_request, + std::vector credentials) ABSL_EXCLUSIVE_LOCKS_REQUIRED(*executor_); absl::flat_hash_map sessions_ ABSL_GUARDED_BY(*executor_); diff --git a/presence/implementation/credential_manager.h b/presence/implementation/credential_manager.h index 84588a26..f26a4a11 100644 --- a/presence/implementation/credential_manager.h +++ b/presence/implementation/credential_manager.h @@ -59,10 +59,15 @@ class CredentialManager { remote_public_creds, UpdateRemotePublicCredentialsCallback credentials_updated_cb) = 0; - // Used to fetch private creds when broadcasting. - virtual void GetPrivateCredentials( + virtual void UpdateLocalCredential( const CredentialSelector& credential_selector, - GetPrivateCredentialsResultCallback callback) = 0; + nearby::internal::LocalCredential credential, + SaveCredentialsResultCallback result_callback) = 0; + + // Used to fetch private creds when broadcasting. + virtual void GetLocalCredentials( + const CredentialSelector& credential_selector, + GetLocalCredentialsResultCallback callback) = 0; // Used to fetch remote public creds when scanning. virtual void GetPublicCredentials( diff --git a/presence/implementation/credential_manager_impl.cc b/presence/implementation/credential_manager_impl.cc index c09bc468..80f1f242 100644 --- a/presence/implementation/credential_manager_impl.cc +++ b/presence/implementation/credential_manager_impl.cc @@ -75,7 +75,7 @@ void CredentialManagerImpl::GenerateCredentials( absl::Duration gap = credential_life_cycle_days * absl::Hours(24); for (int index = 0; index < contiguous_copy_of_credentials; index++) { - auto public_private_credentials = CreatePrivateCredential( + auto public_private_credentials = CreateLocalCredential( device_metadata, identity_type, start_time, start_time + gap); if (public_private_credentials.second.identity_type() != IdentityType::IDENTITY_TYPE_UNSPECIFIED) { @@ -148,7 +148,7 @@ void CredentialManagerImpl::UpdateRemotePublicCredentials( } std::pair -CredentialManagerImpl::CreatePrivateCredential( +CredentialManagerImpl::CreateLocalCredential( const DeviceMetadata& device_metadata, IdentityType identity_type, absl::Time start_time, absl::Time end_time) { LocalCredential private_credential; @@ -301,10 +301,10 @@ std::vector CredentialManagerImpl::ExtendMetadataEncryptionKey( /*info=*/absl::Span(), kNearbyPresenceNumBytesAesGcmKeySize); } -void CredentialManagerImpl::GetPrivateCredentials( +void CredentialManagerImpl::GetLocalCredentials( const CredentialSelector& credential_selector, - GetPrivateCredentialsResultCallback callback) { - credential_storage_ptr_->GetPrivateCredentials(credential_selector, + GetLocalCredentialsResultCallback callback) { + credential_storage_ptr_->GetLocalCredentials(credential_selector, std::move(callback)); } @@ -317,10 +317,10 @@ void CredentialManagerImpl::GetPublicCredentials( } ExceptionOr> -CredentialManagerImpl::GetPrivateCredentialsSync( +CredentialManagerImpl::GetLocalCredentialsSync( const CredentialSelector& credential_selector, absl::Duration timeout) { Future> result; - GetPrivateCredentials( + GetLocalCredentials( credential_selector, {.credentials_fetched_cb = [result](absl::StatusOr> @@ -479,5 +479,13 @@ void CredentialManagerImpl::Subscriber::NotifyCredentialsFetched( callback_.credentials_fetched_cb(credentials); } +void CredentialManagerImpl::UpdateLocalCredential( + const CredentialSelector& credential_selector, + nearby::internal::LocalCredential credential, + SaveCredentialsResultCallback result_callback) { + credential_storage_ptr_->UpdateLocalCredential( + credential_selector.manager_app_id, credential_selector.account_name, + std::move(credential), std::move(result_callback)); +} } // namespace presence } // namespace nearby diff --git a/presence/implementation/credential_manager_impl.h b/presence/implementation/credential_manager_impl.h index af207af2..cf475faf 100644 --- a/presence/implementation/credential_manager_impl.h +++ b/presence/implementation/credential_manager_impl.h @@ -74,13 +74,18 @@ class CredentialManagerImpl : public CredentialManager { remote_public_creds, UpdateRemotePublicCredentialsCallback credentials_updated_cb) override; - void GetPrivateCredentials( + void UpdateLocalCredential( const CredentialSelector& credential_selector, - GetPrivateCredentialsResultCallback callback) override; + nearby::internal::LocalCredential credential, + SaveCredentialsResultCallback result_callback) override; - // Blocking version of `GetPrivateCredentials` + void GetLocalCredentials( + const CredentialSelector& credential_selector, + GetLocalCredentialsResultCallback callback) override; + + // Blocking version of `GetLocalCredentials` nearby::ExceptionOr> - GetPrivateCredentialsSync(const CredentialSelector& credential_selector, + GetLocalCredentialsSync(const CredentialSelector& credential_selector, absl::Duration timeout); // Used to fetch remote public creds when scanning. @@ -109,7 +114,7 @@ class CredentialManagerImpl : public CredentialManager { std::pair - CreatePrivateCredential( + CreateLocalCredential( const nearby::internal::DeviceMetadata& device_metadata, IdentityType identity_type, absl::Time start_time, absl::Time end_time); diff --git a/presence/implementation/credential_manager_impl_test.cc b/presence/implementation/credential_manager_impl_test.cc index 57d035cc..c9ea742a 100644 --- a/presence/implementation/credential_manager_impl_test.cc +++ b/presence/implementation/credential_manager_impl_test.cc @@ -44,8 +44,9 @@ using ::nearby::internal::IdentityType; using ::nearby::internal::LocalCredential; using ::nearby::internal::SharedCredential; using ::nearby::internal::IdentityType::IDENTITY_TYPE_PRIVATE; -using ::proto2::contrib::parse_proto::ParseTestProto; +using ::nearby::internal::IdentityType::IDENTITY_TYPE_TRUSTED; using ::protobuf_matchers::EqualsProto; +using ::testing::UnorderedPointwise; using ::testing::status::StatusIs; constexpr absl::string_view kManagerAppId = "TEST_MANAGER_APP"; @@ -142,7 +143,7 @@ TEST_F(CredentialManagerImplTest, CreateOneCredentialSuccessfully) { constexpr absl::Time kStartTime = absl::FromUnixSeconds(100000); constexpr absl::Time kEndTime = absl::FromUnixSeconds(200000); - auto credentials = credential_manager_.CreatePrivateCredential( + auto credentials = credential_manager_.CreateLocalCredential( device_metadata, IDENTITY_TYPE_PRIVATE, kStartTime, kEndTime); LocalCredential private_credential = credentials.first; @@ -418,7 +419,7 @@ TEST_F(CredentialManagerImplTest, GetLocalCredentialsFailed) { absl::StatusOr> private_credentials; CredentialSelector credential_selector = BuildDefaultCredentialSelector(); - credential_manager_.GetPrivateCredentials( + credential_manager_.GetLocalCredentials( credential_selector, {.credentials_fetched_cb = [&](absl::StatusOr> credentials) { @@ -455,7 +456,7 @@ TEST_F(CredentialManagerImplTest, GetCredentialsSuccessfully) { [&](absl::StatusOr> credentials) { public_credentials = std::move(credentials); }}); - credential_manager_.GetPrivateCredentials( + credential_manager_.GetLocalCredentials( credential_selector, {.credentials_fetched_cb = [&](absl::StatusOr> credentials) { @@ -491,6 +492,58 @@ TEST_F(CredentialManagerImplTest, PublicCredentialsFailEncryption) { EXPECT_THAT(public_credentials, StatusIs(absl::StatusCode::kInvalidArgument)); } +TEST_F(CredentialManagerImplTest, UpdateLocalCredential) { + constexpr int kNumCredentials = 5; + constexpr int kSelectedCredentialId = 2; + constexpr uint16_t kSalt = 1000; + absl::Status update_status = absl::UnknownError(""); + DeviceMetadata device_metadata = CreateTestDeviceMetadata(); + absl::StatusOr> + public_credentials; + std::vector identity_types{IDENTITY_TYPE_PRIVATE, + IDENTITY_TYPE_TRUSTED}; + absl::StatusOr> private_credentials; + absl::StatusOr> modified_private_credentials; + CredentialSelector credential_selector = BuildDefaultCredentialSelector(); + credential_manager_.GenerateCredentials( + device_metadata, kManagerAppId, identity_types, 1, kNumCredentials, + {.credentials_generated_cb = + [&](absl::StatusOr> + credentials) { + public_credentials = std::move(credentials); + }}); + credential_manager_.GetLocalCredentials( + credential_selector, + {.credentials_fetched_cb = + [&](absl::StatusOr> credentials) { + private_credentials = std::move(credentials); + }}); + ASSERT_OK(public_credentials); + ASSERT_OK(private_credentials); + EXPECT_EQ(private_credentials->size(), kNumCredentials); + + // Modify a private credential + LocalCredential& credential = private_credentials->at(kSelectedCredentialId); + credential.mutable_consumed_salts()->insert({kSalt, true}); + + credential_manager_.UpdateLocalCredential( + credential_selector, credential, + {[&](absl::Status status) { update_status = status; }}); + + EXPECT_OK(update_status); + + // verify modified content + credential_manager_.GetLocalCredentials( + credential_selector, + {.credentials_fetched_cb = + [&](absl::StatusOr> credentials) { + modified_private_credentials = std::move(credentials); + }}); + ASSERT_OK(modified_private_credentials); + EXPECT_THAT(*modified_private_credentials, + UnorderedPointwise(EqualsProto(), *private_credentials)); +} + } // namespace } // namespace presence