diff --git a/internal/data/BUILD b/internal/data/BUILD index 9eb465ac..76d0161f 100644 --- a/internal/data/BUILD +++ b/internal/data/BUILD @@ -57,7 +57,6 @@ cc_test( srcs = [ "leveldb_data_set_test.cc", ], - shard_count = 8, deps = [ ":data_manager", ":leveldb_data_set_test_cc_proto", diff --git a/internal/data/data_set.h b/internal/data/data_set.h index 65ce127d..6b8d2474 100644 --- a/internal/data/data_set.h +++ b/internal/data/data_set.h @@ -21,9 +21,9 @@ #include #include "absl/functional/any_invocable.h" +#include "absl/strings/string_view.h" -namespace nearby { -namespace data { +namespace nearby::data { enum class InitStatus { kOK = 0, @@ -42,19 +42,25 @@ class DataSet { virtual ~DataSet() = default; // Asynchronously initializes the object, which must have been created by the - // DataManager::GetDataSet function. |callback| will be invoked on the + // DataManager::GetDataSet function. `callback` will be invoked on the // calling thread when complete. virtual void Initialize(absl::AnyInvocable callback) = 0; - // Asynchronously loads all entries from the database and invokes |callback| + // Asynchronously loads all entries from the database and invokes `callback` // when complete. virtual void LoadEntries( absl::AnyInvocable>) &&> callback) = 0; - // Asynchronously saves |entries_to_save| and deletes entries from - // |keys_to_remove| from the database. |callback| will be invoked on the - // calling thread when complete. |entries_to_save| and |keys_to_remove| must + // Asynchronously loads an entry from the database with key `key` and invokes + // `callback` when complete. + virtual void LoadEntry( + absl::string_view key, + absl::AnyInvocable) &&> callback) = 0; + + // Asynchronously saves `entries_to_save` and deletes entries from + // `keys_to_remove` from the database. `callback` will be invoked on the + // calling thread when complete. `entries_to_save` and `keys_to_remove` must // be non-null. virtual void UpdateEntries( std::unique_ptr entries_to_save, @@ -66,7 +72,6 @@ class DataSet { virtual void Destroy(absl::AnyInvocable callback) = 0; }; -} // namespace data -} // namespace nearby +} // namespace nearby::data #endif // THIRD_PARTY_NEARBY_INTERNAL_DATA_DATA_SET_H_ diff --git a/internal/data/leveldb_data_set.h b/internal/data/leveldb_data_set.h index 5a27efef..c9f625aa 100644 --- a/internal/data/leveldb_data_set.h +++ b/internal/data/leveldb_data_set.h @@ -47,6 +47,9 @@ class LeveldbDataSet : public DataSet { ~LeveldbDataSet() override = default; void Initialize(absl::AnyInvocable callback) override; + void LoadEntry( + absl::string_view key, + absl::AnyInvocable) &&> callback) override; void LoadEntries( absl::AnyInvocable>) &&> callback) override; @@ -84,14 +87,14 @@ void LeveldbDataSet::Initialize( if (status.ok()) { status_ = InitStatus::kOK; - NEARBY_LOGS(INFO) << "Database is initialized successfully.."; + LOG(INFO) << "Database is initialized successfully.."; } else if (status.IsCorruption() || status.IsIOError()) { status_ = InitStatus::kCorrupt; - NEARBY_LOGS(INFO) << "Database is corrupt."; + LOG(INFO) << "Database is corrupt."; } else { status_ = InitStatus::kError; - NEARBY_LOGS(INFO) << "Failed to initialize database due to unknown error."; + LOG(INFO) << "Failed to initialize database due to unknown error."; } std::move(callback)(status_); } @@ -118,16 +121,37 @@ void LeveldbDataSet::LoadEntries( } if (it->status().ok()) { - NEARBY_LOGS(INFO) << "Loaded " << result->size() - << " entries from database."; + LOG(INFO) << "Loaded " << result->size() << " entries from database."; std::move(callback)(true, std::move(result)); } else { - NEARBY_LOGS(INFO) << "Failed to load entries from database."; + LOG(INFO) << "Failed to load entries from database."; result->clear(); std::move(callback)(false, std::move(result)); } } +template ::value, bool> + isMessageLite> +void LeveldbDataSet::LoadEntry( + absl::string_view key, + absl::AnyInvocable) &&> callback) { + auto result = std::make_unique(); + if (status_ != InitStatus::kOK) { + std::move(callback)(false, std::move(result)); + return; + } + + std::string value; + if (!db_->Get(leveldb::ReadOptions(), std::string(key), &value).ok()) { + LOG(INFO) << "Failed to load entry from database with key: " << key; + std::move(callback)(false, std::move(result)); + return; + } + Deserialize(value, *result); + std::move(callback)(true, std::move(result)); +} + template ::value, bool> isMessageLite> @@ -151,11 +175,10 @@ void LeveldbDataSet::LoadEntriesWithKeys( } if (it->status().ok()) { - NEARBY_LOGS(INFO) << "Loaded " << result->size() - << " entries from database."; + LOG(INFO) << "Loaded " << result->size() << " entries from database."; std::move(callback)(true, std::move(result)); } else { - NEARBY_LOGS(INFO) << "Failed to load entries from database."; + LOG(INFO) << "Failed to load entries from database."; result->clear(); std::move(callback)(false, std::move(result)); } @@ -168,7 +191,7 @@ void LeveldbDataSet::UpdateEntries( std::unique_ptr entries_to_save, std::unique_ptr> keys_to_remove, absl::AnyInvocable callback) { - NEARBY_LOGS(INFO) << "UpdateEntries is called."; + LOG(INFO) << "UpdateEntries is called."; if (status_ != InitStatus::kOK) { std::move(callback)(false); return; @@ -196,7 +219,7 @@ template void LeveldbDataSet::Destroy( absl::AnyInvocable callback) { - NEARBY_LOGS(INFO) << "Destroy is called."; + LOG(INFO) << "Destroy is called."; db_.reset(); leveldb::DestroyDB(path_, leveldb::Options()); std::move(callback)(true); diff --git a/internal/data/leveldb_data_set_test.cc b/internal/data/leveldb_data_set_test.cc index c5f2113b..f1843ee2 100644 --- a/internal/data/leveldb_data_set_test.cc +++ b/internal/data/leveldb_data_set_test.cc @@ -34,8 +34,7 @@ #include "internal/data/data_set.h" #include "internal/data/leveldb_data_set_test.proto.h" -namespace nearby { -namespace data { +namespace nearby::data { namespace { using ::testing::SizeIs; @@ -210,6 +209,36 @@ TEST(LeveldbDataSet, LoadEntriesDiceRoll) { EXPECT_EQ((*result)[1].nickname(), "boxcars"); } +TEST(LeveldbDataSet, LoadEntrysDiceRoll) { + std::filesystem::path path = GenerateLeveldbPath(); + std::unique_ptr> diceroll_set = + CreateDataSet(path); + + InitializeAndWait(diceroll_set); + + DiceRoll diceroll1 = GenerateDiceRoll(2); + DiceRoll diceroll2 = GenerateDiceRoll(12); + + auto entries = LeveldbDataSet::KeyEntryVector( + {{"id1", diceroll1}, {"id2", diceroll2}}); + auto data = + std::make_unique::KeyEntryVector>(entries); + UpdateEntriesAndWait(diceroll_set, std::move(data), nullptr); + + absl::Notification notification; + std::unique_ptr result; + diceroll_set->LoadEntry( + "id2", [&result, ¬ification](bool, std::unique_ptr res) { + result = std::move(res); + notification.Notify(); + }); + notification.WaitForNotificationWithTimeout(absl::Seconds(5)); + WipeCleanAndWait(diceroll_set, path); + + EXPECT_THAT(*result, + protobuf_matchers::EqualsProto("value: 12, nickname:'boxcars'")); +} + TEST(LeveldbDataSet, RemoveEntriesDiceRoll) { std::filesystem::path path = GenerateLeveldbPath(); std::unique_ptr> diceroll_set = @@ -255,5 +284,4 @@ TEST(LeveldbDataSet, RemoveEntriesDiceRoll) { } } // namespace -} // namespace data -} // namespace nearby +} // namespace nearby::data diff --git a/sharing/certificates/fake_nearby_share_certificate_storage.cc b/sharing/certificates/fake_nearby_share_certificate_storage.cc index 8df195e4..156da5db 100644 --- a/sharing/certificates/fake_nearby_share_certificate_storage.cc +++ b/sharing/certificates/fake_nearby_share_certificate_storage.cc @@ -14,6 +14,7 @@ #include "sharing/certificates/fake_nearby_share_certificate_storage.h" +#include #include #include #include @@ -29,8 +30,7 @@ #include "sharing/internal/api/public_certificate_database.h" #include "sharing/proto/rpc_resources.pb.h" -namespace nearby { -namespace sharing { +namespace nearby::sharing { using ::nearby::sharing::proto::PublicCertificate; @@ -101,6 +101,14 @@ void FakeNearbyShareCertificateStorage::GetPublicCertificates( get_public_certificates_callbacks_.push_back(std::move(callback)); } +void FakeNearbyShareCertificateStorage::GetPublicCertificate( + absl::string_view id, + std::function< + void(bool, std::unique_ptr)> + callback) { + get_public_certificate_callback_ = std::move(callback); +} + std::optional> FakeNearbyShareCertificateStorage::GetPrivateCertificates() const { return private_certificates_; @@ -152,5 +160,4 @@ void FakeNearbyShareCertificateStorage::SetNextPublicCertificateExpirationTime( next_public_certificate_expiration_time_ = time; } -} // namespace sharing -} // namespace nearby +} // namespace nearby::sharing diff --git a/sharing/certificates/fake_nearby_share_certificate_storage.h b/sharing/certificates/fake_nearby_share_certificate_storage.h index e52326b3..0aa42277 100644 --- a/sharing/certificates/fake_nearby_share_certificate_storage.h +++ b/sharing/certificates/fake_nearby_share_certificate_storage.h @@ -15,6 +15,7 @@ #ifndef THIRD_PARTY_NEARBY_SHARING_CERTIFICATES_FAKE_NEARBY_SHARE_CERTIFICATE_STORAGE_H_ #define THIRD_PARTY_NEARBY_SHARING_CERTIFICATES_FAKE_NEARBY_SHARE_CERTIFICATE_STORAGE_H_ +#include #include #include #include @@ -102,6 +103,11 @@ class FakeNearbyShareCertificateStorage : public NearbyShareCertificateStorage { // NearbyShareCertificateStorage: std::vector GetPublicCertificateIds() const override; void GetPublicCertificates(PublicCertificateCallback callback) override; + void GetPublicCertificate( + absl::string_view id, + std::function)> + callback) override; std::optional> GetPrivateCertificates() const override; std::optional NextPublicCertificateExpirationTime() @@ -153,6 +159,9 @@ class FakeNearbyShareCertificateStorage : public NearbyShareCertificateStorage { std::optional> private_certificates_; std::vector get_public_certificates_callbacks_; + std::function)> + get_public_certificate_callback_; std::vector add_public_certificates_calls_; std::vector remove_expired_public_certificates_calls_; diff --git a/sharing/certificates/nearby_share_certificate_storage.h b/sharing/certificates/nearby_share_certificate_storage.h index c33c6e70..27d4bb0d 100644 --- a/sharing/certificates/nearby_share_certificate_storage.h +++ b/sharing/certificates/nearby_share_certificate_storage.h @@ -21,15 +21,14 @@ #include #include +#include "absl/strings/string_view.h" #include "absl/time/time.h" #include "absl/types/span.h" #include "sharing/certificates/nearby_share_private_certificate.h" -#include "sharing/common/nearby_share_enums.h" #include "sharing/proto/enums.pb.h" #include "sharing/proto/rpc_resources.pb.h" -namespace nearby { -namespace sharing { +namespace nearby::sharing { // Stores local-device private certificates and remote-device public // certificates. Provides methods to help manage certificate expiration. Due to @@ -51,6 +50,13 @@ class NearbyShareCertificateStorage { // Returns all public certificates currently in storage. No RPC call is made. virtual void GetPublicCertificates(PublicCertificateCallback callback) = 0; + // Returns a single public certificate with the given id. + virtual void GetPublicCertificate( + absl::string_view id, + std::function)> + callback) = 0; + // Returns all private certificates currently in storage. Will return // absl::nullopt if deserialization from prefs fails -- not expected to happen // under normal circumstances. @@ -102,7 +108,6 @@ class NearbyShareCertificateStorage { virtual void ClearPublicCertificates(ResultCallback callback) = 0; }; -} // namespace sharing -} // namespace nearby +} // namespace nearby::sharing #endif // THIRD_PARTY_NEARBY_SHARING_CERTIFICATES_NEARBY_SHARE_CERTIFICATE_STORAGE_H_ diff --git a/sharing/certificates/nearby_share_certificate_storage_impl.cc b/sharing/certificates/nearby_share_certificate_storage_impl.cc index 49b90f91..ee2eb3ab 100644 --- a/sharing/certificates/nearby_share_certificate_storage_impl.cc +++ b/sharing/certificates/nearby_share_certificate_storage_impl.cc @@ -44,8 +44,7 @@ #include "sharing/proto/rpc_resources.pb.h" #include "sharing/proto/timestamp.pb.h" -namespace nearby { -namespace sharing { +namespace nearby::sharing { namespace { using ::nearby::sharing::api::PreferenceManager; using ::nearby::sharing::api::PrivateCertificateData; @@ -158,10 +157,10 @@ void NearbyShareCertificateStorageImpl::Initialize() { break; } - NL_VLOG(1) << __func__ - << ": Attempting to initialize public certificate " - "database. Number of attempts: " - << num_initialize_attempts_; + VLOG(1) << __func__ + << ": Attempting to initialize public certificate " + "database. Number of attempts: " + << num_initialize_attempts_; public_certificate_database_->Initialize( [weak_this = weak_from_this()](PublicCertificateDatabase::InitStatus status) { @@ -171,23 +170,23 @@ void NearbyShareCertificateStorageImpl::Initialize() { }); break; case InitStatus::kInitialized: - NL_LOG(INFO) << __func__ << " already initialized."; + LOG(INFO) << __func__ << " already initialized."; break; } } void NearbyShareCertificateStorageImpl::DestroyAndReinitialize() { - NL_LOG(ERROR) << __func__ - << ": Public certificate database corrupt. Erasing and " - "initializing new database."; + LOG(ERROR) << __func__ + << ": Public certificate database corrupt. Erasing and " + "initializing new database."; init_status_ = InitStatus::kUninitialized; public_certificate_database_->Destroy( [weak_this = weak_from_this()](bool success) { if (auto storage = weak_this.lock()) { storage->OnDatabaseDestroyedReinitialize( [&](bool result) { - NL_LOG(INFO) - << "Destroy and reinitialize database. result: " << result; + LOG(INFO) << "Destroy and reinitialize database. result: " + << result; }, success); } @@ -197,8 +196,8 @@ void NearbyShareCertificateStorageImpl::DestroyAndReinitialize() { void NearbyShareCertificateStorageImpl::OnDatabaseInitialized( absl::Time initialize_start_time, PublicCertificateDatabase::InitStatus status) { - NL_LOG(INFO) << "Database is initialized for certificates. status=" - << static_cast(status); + LOG(INFO) << "Database is initialized for certificates. status=" + << static_cast(status); switch (status) { case PublicCertificateDatabase::InitStatus::kOk: FinishInitialization(true); @@ -218,11 +217,11 @@ void NearbyShareCertificateStorageImpl::FinishInitialization(bool success) { // Need to reset the initialize attempts. num_initialize_attempts_ = 0; - NL_VLOG(1) << __func__ - << "Public certificate database initialization succeeded."; + VLOG(1) << __func__ + << "Public certificate database initialization succeeded."; } else { - NL_LOG(ERROR) << __func__ - << "Public certificate database initialization failed."; + LOG(ERROR) << __func__ + << "Public certificate database initialization failed."; } // We run deferred callbacks even if initialization failed not to cause @@ -237,8 +236,8 @@ void NearbyShareCertificateStorageImpl::FinishInitialization(bool success) { void NearbyShareCertificateStorageImpl::OnDatabaseDestroyedReinitialize( ResultCallback callback, bool success) { if (!success) { - NL_LOG(ERROR) << __func__ - << ": Failed to destroy public certificate database."; + LOG(ERROR) << __func__ + << ": Failed to destroy public certificate database."; FinishInitialization(false); callback(false); return; @@ -254,8 +253,8 @@ void NearbyShareCertificateStorageImpl::OnDatabaseDestroyedReinitialize( void NearbyShareCertificateStorageImpl::OnDatabaseDestroyed( ResultCallback callback, bool success) { if (!success) { - NL_LOG(ERROR) << __func__ - << ": Failed to destroy public certificate database."; + LOG(ERROR) << __func__ + << ": Failed to destroy public certificate database."; std::move(callback)(false); return; } @@ -270,11 +269,11 @@ void NearbyShareCertificateStorageImpl::AddPublicCertificatesCallback( std::unique_ptr new_expirations, ResultCallback callback, bool proceed) { if (!proceed) { - NL_LOG(ERROR) << __func__ << ": Failed to add public certificates."; + LOG(ERROR) << __func__ << ": Failed to add public certificates."; std::move(callback)(false); return; } - NL_VLOG(1) << __func__ << ": Successfully added public certificates."; + VLOG(1) << __func__ << ": Successfully added public certificates."; public_certificate_expirations_ = MergeExpirations(public_certificate_expirations_, *new_expirations); @@ -286,13 +285,11 @@ void NearbyShareCertificateStorageImpl::RemoveExpiredPublicCertificatesCallback( const absl::flat_hash_set& ids_to_remove, ResultCallback callback, bool proceed) { if (!proceed) { - NL_LOG(ERROR) << __func__ - << ": Failed to remove expired public certificates."; + LOG(ERROR) << __func__ << ": Failed to remove expired public certificates."; std::move(callback)(false); return; } - NL_VLOG(1) << __func__ - << ": Expired public certificates successfully removed."; + VLOG(1) << __func__ << ": Expired public certificates successfully removed."; auto should_remove = [&](const std::pair& pair) -> bool { @@ -330,10 +327,31 @@ void NearbyShareCertificateStorageImpl::GetPublicCertificates( return; } - NL_VLOG(1) << __func__ << ": Calling LoadEntries on database."; + VLOG(1) << __func__ << ": Calling LoadEntries on database."; public_certificate_database_->LoadEntries(std::move(callback)); } +void NearbyShareCertificateStorageImpl::GetPublicCertificate( + absl::string_view id, + std::function< + void(bool, std::unique_ptr)> + callback) { + if (init_status_ == InitStatus::kFailed) { + std::move(callback)(false, nullptr); + return; + } + + if (init_status_ == InitStatus::kUninitialized) { + deferred_callbacks_.push( + [this, id = std::string(id), callback = std::move(callback)]() mutable { + GetPublicCertificate(id, std::move(callback)); + }); + return; + } + VLOG(1) << __func__ << ": Calling LoadCertificate on database, key: " << id; + public_certificate_database_->LoadCertificate(id, std::move(callback)); +} + std::optional> NearbyShareCertificateStorageImpl::GetPrivateCertificates() const { std::vector list = @@ -395,9 +413,9 @@ void NearbyShareCertificateStorageImpl::AddPublicCertificates( } std::sort(new_expirations.begin(), new_expirations.end(), SortBySecond); - NL_VLOG(1) << __func__ - << ": Calling UpdateEntries on public certificate database with " - << public_certificates.size() << " certificates."; + VLOG(1) << __func__ + << ": Calling UpdateEntries on public certificate database with " + << public_certificates.size() << " certificates."; public_certificate_database_->AddCertificates( public_certificates, [weak_this = weak_from_this(), new_expirations, callback = std::move(callback)](bool success) { @@ -443,10 +461,9 @@ void NearbyShareCertificateStorageImpl::RemoveExpiredPublicCertificates( return; } - NL_VLOG(1) - << __func__ - << ": Calling UpdateEntries on public certificate database to remove " - << ids_to_remove.size() << " expired certificates."; + VLOG(1) << __func__ + << ": Calling UpdateEntries on public certificate database to remove " + << ids_to_remove.size() << " expired certificates."; absl::flat_hash_set remove_set(ids_to_remove.begin(), ids_to_remove.end()); public_certificate_database_->RemoveCertificatesById( @@ -467,7 +484,7 @@ void NearbyShareCertificateStorageImpl::ClearPublicCertificates( return; } - NL_VLOG(1) << __func__ << ": Calling Destroy on public certificate database."; + VLOG(1) << __func__ << ": Calling Destroy on public certificate database."; init_status_ = InitStatus::kUninitialized; public_certificate_database_->Destroy( [weak_this = weak_from_this(), @@ -515,5 +532,4 @@ void NearbyShareCertificateStorageImpl::SavePublicCertificateExpirations() { prefs::kNearbySharingPublicCertificateExpirationDictName, expirations); } -} // namespace sharing -} // namespace nearby +} // namespace nearby::sharing diff --git a/sharing/certificates/nearby_share_certificate_storage_impl.h b/sharing/certificates/nearby_share_certificate_storage_impl.h index aa623c5b..256aafc0 100644 --- a/sharing/certificates/nearby_share_certificate_storage_impl.h +++ b/sharing/certificates/nearby_share_certificate_storage_impl.h @@ -34,8 +34,7 @@ #include "sharing/internal/api/public_certificate_database.h" #include "sharing/proto/rpc_resources.pb.h" -namespace nearby { -namespace sharing { +namespace nearby::sharing { // Implements NearbyShareCertificateStorage using Prefs to store private // certificates and LevelDB Proto to store public certificates. Must be @@ -74,6 +73,11 @@ class NearbyShareCertificateStorageImpl : public NearbyShareCertificateStorage, // NearbyShareCertificateStorage std::vector GetPublicCertificateIds() const override; void GetPublicCertificates(PublicCertificateCallback callback) override; + void GetPublicCertificate( + absl::string_view id, + std::function)> + callback) override; std::optional> GetPrivateCertificates() const override; std::optional NextPublicCertificateExpirationTime() @@ -128,7 +132,6 @@ class NearbyShareCertificateStorageImpl : public NearbyShareCertificateStorage, std::queue> deferred_callbacks_; }; -} // namespace sharing -} // namespace nearby +} // namespace nearby::sharing #endif // THIRD_PARTY_NEARBY_SHARING_CERTIFICATES_NEARBY_SHARE_CERTIFICATE_STORAGE_IMPL_H_ diff --git a/sharing/certificates/nearby_share_certificate_storage_impl_test.cc b/sharing/certificates/nearby_share_certificate_storage_impl_test.cc index 54c779fe..e000d98f 100644 --- a/sharing/certificates/nearby_share_certificate_storage_impl_test.cc +++ b/sharing/certificates/nearby_share_certificate_storage_impl_test.cc @@ -46,8 +46,7 @@ #include "sharing/proto/rpc_resources.pb.h" #include "sharing/proto/timestamp.pb.h" -namespace nearby { -namespace sharing { +namespace nearby::sharing { namespace { using ::nearby::sharing::api::MockPublicCertificateDb; using ::nearby::sharing::proto::DeviceVisibility; @@ -428,6 +427,32 @@ TEST_F(NearbyShareCertificateStorageImplTest, GetPublicCertificates) { EXPECT_THAT(cert_store.use_count(), Eq(1)); } +TEST_F(NearbyShareCertificateStorageImplTest, GetPublicCertificate) { + auto db = std::make_unique( + PrepopulatePublicCertificates()); + nearby::FakePublicCertificateDb* fake_db = db.get(); + + auto cert_store = NearbyShareCertificateStorageImpl::Factory::Create( + preference_manager_, std::move(db)); + fake_db->InvokeInitStatusCallback(FakePublicCertificateDb::InitStatus::kOk); + + std::unique_ptr public_certificate; + cert_store->GetPublicCertificate( + kSecretId3, [&public_certificate]( + bool success, std::unique_ptr result) { + public_certificate = std::move(result); + }); + fake_db->InvokeLoadCertificateCallback(true); + + std::string expected_serialized, actual_serialized; + ASSERT_TRUE(public_certificate->SerializeToString(&actual_serialized)); + ASSERT_TRUE(fake_db->GetCertificatesMap() + .find(kSecretId3) + ->second.SerializeToString(&expected_serialized)); + ASSERT_EQ(expected_serialized, actual_serialized); + EXPECT_THAT(cert_store.use_count(), Eq(1)); +} + TEST_F(NearbyShareCertificateStorageImplTest, AddPublicCertificates) { auto db = std::make_unique( PrepopulatePublicCertificates()); @@ -784,5 +809,4 @@ TEST_F(NearbyShareCertificateStorageImplTest, EXPECT_THAT(cert_store.use_count(), Eq(1)); } -} // namespace sharing -} // namespace nearby +} // namespace nearby::sharing diff --git a/sharing/internal/api/mock_public_certificate_db.h b/sharing/internal/api/mock_public_certificate_db.h index e289ce54..61f618dd 100644 --- a/sharing/internal/api/mock_public_certificate_db.h +++ b/sharing/internal/api/mock_public_certificate_db.h @@ -21,6 +21,7 @@ #include "gmock/gmock.h" #include "absl/functional/any_invocable.h" +#include "absl/strings/string_view.h" #include "absl/types/span.h" #include "sharing/internal/api/public_certificate_database.h" @@ -41,6 +42,13 @@ class MockPublicCertificateDb : public PublicCertificateDatabase { nearby::sharing::proto::PublicCertificate>>) &&> callback), (override)); + MOCK_METHOD( + void, LoadCertificate, + (absl::string_view id, + absl::AnyInvocable) &&> + callback), + (override)); MOCK_METHOD( void, AddCertificates, (absl::Span certificates, diff --git a/sharing/internal/api/public_certificate_database.h b/sharing/internal/api/public_certificate_database.h index 4c6b0e51..5c50f298 100644 --- a/sharing/internal/api/public_certificate_database.h +++ b/sharing/internal/api/public_certificate_database.h @@ -20,6 +20,7 @@ #include #include "absl/functional/any_invocable.h" +#include "absl/strings/string_view.h" #include "absl/types/span.h" #include "sharing/proto/rpc_resources.pb.h" @@ -53,6 +54,13 @@ class PublicCertificateDatabase { nearby::sharing::proto::PublicCertificate>>) &&> callback) = 0; + virtual void LoadCertificate( + absl::string_view id, + absl::AnyInvocable< + void(bool, + std::unique_ptr) &&> + callback) = 0; + // Asynchronously saves |certificates| to the database. // |callback| can be invoked on an executor thread when complete. virtual void AddCertificates( diff --git a/sharing/internal/test/fake_public_certificate_db.cc b/sharing/internal/test/fake_public_certificate_db.cc index 1d19451e..0d9db02b 100644 --- a/sharing/internal/test/fake_public_certificate_db.cc +++ b/sharing/internal/test/fake_public_certificate_db.cc @@ -21,6 +21,7 @@ #include #include "absl/functional/any_invocable.h" +#include "absl/strings/string_view.h" #include "absl/types/span.h" #include "sharing/internal/api/public_certificate_database.h" #include "sharing/proto/rpc_resources.pb.h" @@ -46,6 +47,14 @@ void FakePublicCertificateDb::LoadEntries( load_callback_ = std::move(callback); } +void FakePublicCertificateDb::LoadCertificate( + absl::string_view id, + absl::AnyInvocable) &&> + callback) { + load_certificate_id_ = id; + load_certificate_callback_ = std::move(callback); +} + void FakePublicCertificateDb::AddCertificates( absl::Span certificates, absl::AnyInvocable callback) { @@ -90,6 +99,16 @@ void FakePublicCertificateDb::InvokeLoadCallback(bool success) { std::move(load_callback_)(success, std::move(result)); } +void FakePublicCertificateDb::InvokeLoadCertificateCallback(bool success) { + const auto& it = entries_.find(load_certificate_id_); + if (it == entries_.end()) { + std::move(load_certificate_callback_)(success, nullptr); + return; + } + std::move(load_certificate_callback_)( + success, std::make_unique(it->second)); +} + void FakePublicCertificateDb::InvokeAddCallback(bool success) { std::move(add_callback_)(success); } diff --git a/sharing/internal/test/fake_public_certificate_db.h b/sharing/internal/test/fake_public_certificate_db.h index 6cb417cd..51fa46dc 100644 --- a/sharing/internal/test/fake_public_certificate_db.h +++ b/sharing/internal/test/fake_public_certificate_db.h @@ -21,6 +21,7 @@ #include #include "absl/functional/any_invocable.h" +#include "absl/strings/string_view.h" #include "absl/types/span.h" #include "sharing/internal/api/public_certificate_database.h" @@ -42,6 +43,12 @@ class FakePublicCertificateDb void(bool, std::unique_ptr>) &&> callback) override; + void LoadCertificate( + absl::string_view id, + absl::AnyInvocable< + void(bool, std::unique_ptr) + &&> + callback) override; void AddCertificates( absl::Span certificates, absl::AnyInvocable callback) override; @@ -59,6 +66,7 @@ class FakePublicCertificateDb void InvokeInitStatusCallback( nearby::sharing::api::PublicCertificateDatabase::InitStatus init_status); void InvokeLoadCallback(bool success); + void InvokeLoadCertificateCallback(bool success); void InvokeAddCallback(bool success); void InvokeRemoveCallback(bool success); void InvokeDestroyCallback(bool success); @@ -73,6 +81,10 @@ class FakePublicCertificateDb void(bool, std::unique_ptr>) &&> load_callback_; + std::string load_certificate_id_; + absl::AnyInvocable< + void(bool, std::unique_ptr) &&> + load_certificate_callback_; absl::AnyInvocable add_callback_; absl::AnyInvocable remove_callback_; absl::AnyInvocable destroy_callback_;