Store consumed salts in private credentials

PiperOrigin-RevId: 505167378
This commit is contained in:
Janusz Sobczak
2023-01-27 11:30:54 -08:00
committed by Copybara-Service
parent 20989c5167
commit 94e70d0ddd
16 changed files with 389 additions and 132 deletions
+11 -4
View File
@@ -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(
+7 -2
View File
@@ -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(
@@ -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<LocalCredential> BuildPrivateCreds(absl::string_view secret_id) {
std::vector<LocalCredential> 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<SharedCredential> BuildPublicCreds(absl::string_view secret_id) {
return public_credentials;
}
absl::StatusOr<std::vector<LocalCredential>> GetPrivateCredentials(
absl::StatusOr<std::vector<LocalCredential>> 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<std::vector<LocalCredential>> GetPrivateCredentials(
.account_name = std::string(account_name),
.identity_type = identity_type};
absl::StatusOr<std::vector<LocalCredential>> private_credentials;
credential_storage.GetPrivateCredentials(
credential_storage.GetLocalCredentials(
selector,
GetPrivateCredentialsResultCallback{
GetLocalCredentialsResultCallback{
.credentials_fetched_cb =
[&](absl::StatusOr<std::vector<LocalCredential>> 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<SharedCredential>(),
@@ -162,7 +162,7 @@ absl::Status SavePublicCredentials(CredentialStorageImpl& credential_storage,
BuildPublicCreds(secret_id), credential_type);
}
TEST(CredentialStorageImplTest, SaveAndGetPrivateCredentials) {
TEST(CredentialStorageImplTest, SaveAndGetLocalCredentials) {
std::vector<LocalCredential> default_private_creds =
BuildPrivateCreds(kSecretId);
std::vector<SharedCredential> 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<LocalCredential> default_private_creds =
BuildPrivateCreds(kSecretId);
std::vector<SharedCredential> 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<IdentityType> {};
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<LocalCredential>{CreatePrivateCredential(
kSecretId, identity_type)}));
std::vector<LocalCredential>{
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<LocalCredential> 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));
@@ -74,7 +74,7 @@ struct UpdateRemotePublicCredentialsCallback {
absl::AnyInvocable<void(absl::Status)> credentials_updated_cb;
};
struct GetPrivateCredentialsResultCallback {
struct GetLocalCredentialsResultCallback {
absl::AnyInvocable<void(
absl::StatusOr<std::vector<nearby::internal::LocalCredential>>)>
credentials_fetched_cb;
@@ -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.
//
@@ -14,6 +14,7 @@
#include "internal/platform/implementation/g3/credential_storage_impl.h"
#include <algorithm>
#include <string>
#include <tuple>
#include <utility>
@@ -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<LocalCredential>& 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<std::vector<LocalCredential>> 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<LocalCredential>();
}
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<std::vector<nearby::internal::LocalCredential>>
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<LocalCredential> 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
@@ -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<std::string, std::string>;
using LocalCredentialKey = std::pair<std::string, std::string>;
using PublicCredentialKey =
std::tuple<std::string, std::string, PublicCredentialType>;
@@ -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<PrivateCredentialKey, std::vector<LocalCredential>>
absl::StatusOr<std::vector<LocalCredential>> GetLocalCredentialsLocked(
const CredentialSelector& credential_selector);
void SaveLocalCredentialsLocked(
absl::string_view manager_app_id, absl::string_view account_name,
const std::vector<LocalCredential>& private_credentials);
absl::flat_hash_map<LocalCredentialKey, std::vector<LocalCredential>>
private_credentials_map_;
absl::flat_hash_map<PublicCredentialKey, std::vector<SharedCredential>>
public_credentials_map_;
@@ -15,6 +15,7 @@
#include "presence/implementation/advertisement_factory.h"
#include <string>
#include <utility>
#include <vector>
#include "absl/status/status.h"
@@ -96,11 +97,11 @@ std::string SerializeAction(const Action& action) {
absl::StatusOr<AdvertisementData> AdvertisementFactory::CreateAdvertisement(
const BaseBroadcastRequest& request,
std::vector<LocalCredential>& credentials) const {
absl::optional<LocalCredential> credential) const {
AdvertisementData advert = {};
if (absl::holds_alternative<BaseBroadcastRequest::BasePresence>(
request.variant)) {
return CreateBaseNpAdvertisement(request, credentials);
return CreateBaseNpAdvertisement(request, std::move(credential));
}
return advert;
}
@@ -108,7 +109,7 @@ absl::StatusOr<AdvertisementData> AdvertisementFactory::CreateAdvertisement(
absl::StatusOr<AdvertisementData>
AdvertisementFactory::CreateBaseNpAdvertisement(
const BaseBroadcastRequest& request,
std::vector<LocalCredential>& credentials) const {
absl::optional<LocalCredential> credential) const {
const auto& presence =
absl::get<BaseBroadcastRequest::BasePresence>(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<std::string> 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<std::string> AdvertisementFactory::EncryptDataElements(
std::vector<LocalCredential>& 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",
@@ -39,20 +39,19 @@ class AdvertisementFactory {
// Returns a BLE advertisement for given `request.
absl::StatusOr<AdvertisementData> CreateAdvertisement(
const BaseBroadcastRequest& request,
std::vector<LocalCredential>& credentials) const;
absl::optional<LocalCredential> credential) const;
absl::StatusOr<AdvertisementData> CreateAdvertisement(
const BaseBroadcastRequest& request) const {
std::vector<LocalCredential> empty;
return CreateAdvertisement(request, empty);
return CreateAdvertisement(request, absl::optional<LocalCredential>());
}
private:
absl::StatusOr<AdvertisementData> CreateBaseNpAdvertisement(
const BaseBroadcastRequest& request,
std::vector<LocalCredential>& credentials) const;
absl::optional<LocalCredential> credential) const;
absl::StatusOr<std::string> EncryptDataElements(
std::vector<LocalCredential>& credentials, absl::string_view salt,
const LocalCredential& credential, absl::string_view salt,
absl::string_view data_elements) const;
};
@@ -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<LocalCredential> credentials = {
CreatePrivateCredential(kIdentity)};
std::vector<DataElement> 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<AdvertisementData> 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<LocalCredential> credentials = {
CreatePrivateCredential(kIdentity)};
std::vector<DataElement> 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<AdvertisementData> 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<LocalCredential> credentials = {
CreatePrivateCredential(kIdentity)};
std::vector<DataElement> 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<AdvertisementData> result =
AdvertisementFactory().CreateAdvertisement(request, credentials);
AdvertisementFactory().CreateAdvertisement(
request, CreateLocalCredential(kIdentity));
ASSERT_OK(result);
EXPECT_FALSE(result->is_extended_advertisement);
+97 -14
View File
@@ -14,11 +14,18 @@
#include "presence/implementation/broadcast_manager.h"
#include <algorithm>
#include <limits>
#include <memory>
#include <set>
#include <string>
#include <utility>
#include <vector>
#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<uint16_t>();
}
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<LocalCredential> 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<LocalCredential> credentials) {
absl::optional<LocalCredential> BroadcastManager::SelectCredential(
BaseBroadcastRequest& broadcast_request,
std::vector<LocalCredential> credentials) {
if (credentials.empty()) {
return absl::optional<LocalCredential>();
}
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<LocalCredential>();
}
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<LocalCredential> BroadcastManager::Advertise(
BroadcastSessionId id, BaseBroadcastRequest broadcast_request,
std::vector<LocalCredential> credentials) {
auto it = sessions_.find(id);
if (it == sessions_.end()) {
NEARBY_LOGS(INFO) << "Broadcast session terminated, id: " << id;
return;
return absl::optional<LocalCredential>();
}
absl::optional<LocalCredential> credential =
SelectCredential(broadcast_request, std::move(credentials));
absl::StatusOr<AdvertisementData> 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<LocalCredential>();
}
std::unique_ptr<AdvertisingSession> 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<LocalCredential>();
}
it->second.SetAdvertisingSession(std::move(session));
return credential;
}
void BroadcastManager::NotifyStartCallbackStatus(BroadcastSessionId id,
+10 -2
View File
@@ -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<LocalCredential> SelectCredential(
BaseBroadcastRequest& broadcast_request,
std::vector<LocalCredential> credentials);
void Advertise(BroadcastSessionId id, BaseBroadcastRequest broadcast_request,
std::vector<LocalCredential> 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<LocalCredential> Advertise(
BroadcastSessionId id, BaseBroadcastRequest broadcast_request,
std::vector<LocalCredential> credentials)
ABSL_EXCLUSIVE_LOCKS_REQUIRED(*executor_);
absl::flat_hash_map<BroadcastSessionId, BroadcastSessionState> sessions_
ABSL_GUARDED_BY(*executor_);
+8 -3
View File
@@ -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(
@@ -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<LocalCredential, SharedCredential>
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<uint8_t> CredentialManagerImpl::ExtendMetadataEncryptionKey(
/*info=*/absl::Span<uint8_t>(), 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<std::vector<LocalCredential>>
CredentialManagerImpl::GetPrivateCredentialsSync(
CredentialManagerImpl::GetLocalCredentialsSync(
const CredentialSelector& credential_selector, absl::Duration timeout) {
Future<std::vector<LocalCredential>> result;
GetPrivateCredentials(
GetLocalCredentials(
credential_selector,
{.credentials_fetched_cb =
[result](absl::StatusOr<std::vector<LocalCredential>>
@@ -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
@@ -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<std::vector<nearby::internal::LocalCredential>>
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<nearby::internal::LocalCredential,
nearby::internal::SharedCredential>
CreatePrivateCredential(
CreateLocalCredential(
const nearby::internal::DeviceMetadata& device_metadata,
IdentityType identity_type, absl::Time start_time, absl::Time end_time);
@@ -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<std::vector<LocalCredential>> private_credentials;
CredentialSelector credential_selector = BuildDefaultCredentialSelector();
credential_manager_.GetPrivateCredentials(
credential_manager_.GetLocalCredentials(
credential_selector,
{.credentials_fetched_cb =
[&](absl::StatusOr<std::vector<LocalCredential>> credentials) {
@@ -455,7 +456,7 @@ TEST_F(CredentialManagerImplTest, GetCredentialsSuccessfully) {
[&](absl::StatusOr<std::vector<SharedCredential>> credentials) {
public_credentials = std::move(credentials);
}});
credential_manager_.GetPrivateCredentials(
credential_manager_.GetLocalCredentials(
credential_selector,
{.credentials_fetched_cb =
[&](absl::StatusOr<std::vector<LocalCredential>> 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<std::vector<nearby::internal::SharedCredential>>
public_credentials;
std::vector<IdentityType> identity_types{IDENTITY_TYPE_PRIVATE,
IDENTITY_TYPE_TRUSTED};
absl::StatusOr<std::vector<LocalCredential>> private_credentials;
absl::StatusOr<std::vector<LocalCredential>> modified_private_credentials;
CredentialSelector credential_selector = BuildDefaultCredentialSelector();
credential_manager_.GenerateCredentials(
device_metadata, kManagerAppId, identity_types, 1, kNumCredentials,
{.credentials_generated_cb =
[&](absl::StatusOr<std::vector<nearby::internal::SharedCredential>>
credentials) {
public_credentials = std::move(credentials);
}});
credential_manager_.GetLocalCredentials(
credential_selector,
{.credentials_fetched_cb =
[&](absl::StatusOr<std::vector<LocalCredential>> 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<std::vector<LocalCredential>> credentials) {
modified_private_credentials = std::move(credentials);
}});
ASSERT_OK(modified_private_credentials);
EXPECT_THAT(*modified_private_credentials,
UnorderedPointwise(EqualsProto(), *private_credentials));
}
} // namespace
} // namespace presence