blob: a34894f29df41f0462be203b11dd4ba87d1f4ad7 [file]
// Copyright 2024 The Chromium Authors
// Use of this source code is governed by a BSD-style license that can be
// found in the LICENSE file.
#include "net/device_bound_sessions/session_store_impl.h"
#include <algorithm>
#include <optional>
#include "base/containers/map_util.h"
#include "base/containers/span.h"
#include "base/debug/dump_without_crashing.h"
#include "base/metrics/histogram_functions.h"
#include "base/process/process.h"
#include "base/sequence_checker.h"
#include "base/strings/string_view_util.h"
#include "base/task/sequenced_task_runner.h"
#include "base/task/thread_pool.h"
#include "base/time/time.h"
#include "base/types/expected_macros.h"
#include "components/unexportable_keys/background_task_priority.h"
#include "components/unexportable_keys/features.h"
#include "components/unexportable_keys/service_error.h"
#include "components/unexportable_keys/unexportable_key_id.h"
#include "components/unexportable_keys/unexportable_key_service.h"
#include "net/base/features.h"
#include "net/base/schemeful_site.h"
#include "net/device_bound_sessions/deletion_reason.h"
#include "net/device_bound_sessions/proto/storage.pb.h"
namespace net::device_bound_sessions {
namespace {
using unexportable_keys::BackgroundTaskPriority;
using unexportable_keys::ServiceError;
using unexportable_keys::ServiceErrorOr;
using unexportable_keys::UnexportableKeyService;
using unexportable_keys::UnexportableSigningKeyId;
// Priority is set to `USER_VISIBLE` because the initial load of
// sessions from disk is required to complete before URL requests
// can be checked to see if they are associated with bound sessions.
constexpr base::TaskTraits kDBTaskTraits = {
base::MayBlock(), base::TaskPriority::USER_VISIBLE,
base::TaskShutdownBehavior::BLOCK_SHUTDOWN};
const char kSessionTableName[] = "dbsc_session_tbl";
const base::TimeDelta kFlushDelay = base::Seconds(2);
// The delay between when the session service is loaded and the garbage
// collection is started. This is delayed to not slow down the startup of the
// browser.
constexpr base::TimeDelta kGarbageCollectionDelay = base::Minutes(2);
// Histogram name for the garbage collection of unexportable keys.
constexpr std::string_view kGarbageCollectionHistogramPrefix =
"Crypto.UnexportableKeys.GarbageCollection.DeviceBoundSessions.";
SessionStoreImpl::DBStatus InitializeOnDbSequence(
sql::Database* db,
base::FilePath db_storage_path,
sqlite_proto::ProtoTableManager* table_manager,
sqlite_proto::KeyValueData<proto::SiteSessions>* session_data) {
if (db->Open(db_storage_path) == false) {
return SessionStoreImpl::DBStatus::kFailure;
}
// Control the schema version with a feature param so that the database can be
// wiped between Origin Trials and going into the final release.
table_manager->InitializeOnDbSequence(
db, std::vector<std::string>{kSessionTableName},
features::kDeviceBoundSessionsSchemaVersion.Get());
session_data->InitializeOnDBSequence();
return SessionStoreImpl::DBStatus::kSuccess;
}
} // namespace
SessionStoreImpl::SessionStoreImpl(base::FilePath db_storage_path,
UnexportableKeyService& key_service)
: key_service_(key_service),
db_task_runner_(
base::ThreadPool::CreateSequencedTaskRunner(kDBTaskTraits)),
db_storage_path_(std::move(db_storage_path)),
db_(std::make_unique<sql::Database>(sql::DatabaseOptions(),
sql::Database::Tag("DBSCSessions"))),
table_manager_(base::MakeRefCounted<sqlite_proto::ProtoTableManager>(
db_task_runner_)),
session_table_(
std::make_unique<sqlite_proto::KeyValueTable<proto::SiteSessions>>(
kSessionTableName)),
session_data_(
std::make_unique<sqlite_proto::KeyValueData<proto::SiteSessions>>(
table_manager_,
session_table_.get(),
/*max_num_entries=*/std::nullopt,
kFlushDelay)) {}
SessionStoreImpl::~SessionStoreImpl() {
DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
if (db_status_ == DBStatus::kSuccess) {
session_data_->FlushDataToDisk();
}
// Shutdown `table_manager_`, and delete it together with `db_`
// and KeyValueTable on DB sequence, then delete the KeyValueData
// and call `shutdown_callback_` on main sequence.
// This ensures that DB objects outlive any other task posted to DB
// sequence, since their deletion is the very last posted task.
db_task_runner_->PostTaskAndReply(
FROM_HERE,
base::BindOnce(
[](scoped_refptr<sqlite_proto::ProtoTableManager> table_manager,
std::unique_ptr<sql::Database> db,
auto session_table) { table_manager->WillShutdown(); },
std::move(table_manager_), std::move(db_), std::move(session_table_)),
base::BindOnce(
[](auto session_data, base::OnceClosure shutdown_callback) {
if (shutdown_callback) {
std::move(shutdown_callback).Run();
}
},
std::move(session_data_), std::move(shutdown_callback_)));
}
void SessionStoreImpl::LoadSessions(LoadSessionsCallback callback) {
CHECK_EQ(db_status_, DBStatus::kNotLoaded);
// This is safe because tasks are serialized on the db_task_runner sequence
// and the `table_manager_` and `session_data_` are only freed after a
// response from a task (triggered by the destructor) runs on the
// `db_task_runner_`.
// Similarly, the `db_` is not actually destroyed until the task
// triggered by the destructor runs on the `db_task_runner_`.
db_task_runner_->PostTaskAndReplyWithResult(
FROM_HERE,
base::BindOnce(&InitializeOnDbSequence, base::Unretained(db_.get()),
db_storage_path_, base::Unretained(table_manager_.get()),
base::Unretained(session_data_.get())),
base::BindOnce(&SessionStoreImpl::OnDatabaseLoaded,
weak_ptr_factory_.GetWeakPtr(), std::move(callback),
base::ElapsedTimer()));
}
void SessionStoreImpl::OnDatabaseLoaded(LoadSessionsCallback callback,
base::ElapsedTimer timer,
DBStatus db_status) {
db_status_ = db_status;
SessionsMap sessions;
if (db_status == DBStatus::kSuccess) {
std::vector<std::string> keys_to_delete;
std::map<std::string, proto::SiteSessions> sites_to_update;
sessions = CreateSessionsFromLoadedData(session_data_->GetAllCached(),
keys_to_delete, sites_to_update,
/*prune_expired_sessions=*/true);
if (!keys_to_delete.empty()) {
session_data_->DeleteData(keys_to_delete);
}
for (const auto& [site_str, site_proto] : sites_to_update) {
session_data_->UpdateData(site_str, site_proto);
}
// Schedule a task for original profiles to obtain all keys that were
// created for this profile in the past, including all OTR profiles.
if (base::FeatureList::IsEnabled(
unexportable_keys::kUnexportableKeyDeletion)) {
base::SequencedTaskRunner::GetCurrentDefault()->PostDelayedTask(
FROM_HERE,
base::BindOnce(&SessionStoreImpl::StartGarbageCollection,
weak_ptr_factory_.GetWeakPtr()),
kGarbageCollectionDelay);
}
}
base::UmaHistogramBoolean("Net.DeviceBoundSessions.SessionStoreLoadSuccess",
db_status == DBStatus::kSuccess);
base::UmaHistogramTimes("Net.DeviceBoundSessions.SessionStoreLoadDuration",
timer.Elapsed());
std::move(callback).Run(std::move(sessions));
}
// static
SessionStore::SessionsMap SessionStoreImpl::CreateSessionsFromLoadedData(
const std::map<std::string, proto::SiteSessions>& loaded_data,
std::vector<std::string>& keys_to_delete,
std::map<std::string, proto::SiteSessions>& sites_to_update,
bool prune_expired_sessions) {
SessionsMap all_sessions;
for (const auto& [site_str, site_proto] : loaded_data) {
SchemefulSite site = net::SchemefulSite::Deserialize(site_str);
if (site.opaque()) {
keys_to_delete.push_back(site_str);
continue;
}
SessionsMap site_sessions;
std::vector<std::string> session_ids_to_prune;
for (const auto& [session_id, session_proto] : site_proto.sessions()) {
auto session_or_error = Session::CreateFromProto(
session_proto, /*check_expiry=*/prune_expired_sessions);
if (!session_or_error.has_value()) {
LogSessionDeletionReason(session_or_error.error());
session_ids_to_prune.push_back(session_id);
continue;
}
std::unique_ptr<Session> session = std::move(session_or_error.value());
if (session->id().value() != session_id) {
// TODO(crbug.com/552483536): Replace with session pruning once we
// verify whether this discrepancy occurs in the wild.
base::debug::DumpWithoutCrashing();
}
// Session is structurally valid and unexpired.
site_sessions.emplace(SessionKey{site, session->id()},
std::move(session));
}
// If no valid sessions remain for this site, remove the entire site entry
// from the DB. Otherwise, update the DB entry if some sessions were pruned.
if (site_sessions.empty()) {
keys_to_delete.push_back(site_str);
} else {
if (!session_ids_to_prune.empty()) {
proto::SiteSessions updated_site_proto = site_proto;
for (const std::string& invalid_id : session_ids_to_prune) {
updated_site_proto.mutable_sessions()->erase(invalid_id);
}
sites_to_update[site_str] = std::move(updated_site_proto);
}
all_sessions.merge(site_sessions);
}
}
return all_sessions;
}
void SessionStoreImpl::SetShutdownCallbackForTesting(
base::OnceClosure shutdown_callback) {
shutdown_callback_ = std::move(shutdown_callback);
}
void SessionStoreImpl::SaveSession(const SchemefulSite& site,
const Session& session,
SessionStore::SaveSessionMode mode) {
if (db_status_ != DBStatus::kSuccess) {
return;
}
CHECK(session.unexportable_key_id().has_value());
// Wrap the unexportable key into a persistable form.
ServiceErrorOr<std::vector<uint8_t>> wrapped_key =
key_service_->GetWrappedKey(*session.unexportable_key_id());
// Don't bother persisting the session if wrapping fails because we will throw
// away all persisted data if the wrapped key is missing for any session.
if (!wrapped_key.has_value()) {
return;
}
proto::Session session_proto = session.ToProto();
session_proto.set_wrapped_key(
std::string(wrapped_key->begin(), wrapped_key->end()));
// Handle attestation key if present.
AttestationKeySaveOutcome outcome =
SetWrappedAttestationKey(site, session, session_proto, mode);
base::UmaHistogramEnumeration(
"Net.DeviceBoundSessions.AttestationKeySaveOutcome", outcome);
proto::SiteSessions site_proto;
std::string site_str = site.Serialize();
session_data_->TryGetData(site_str, &site_proto);
(*site_proto.mutable_sessions())[session_proto.id()] =
std::move(session_proto);
session_data_->UpdateData(site_str, site_proto);
}
SessionStoreImpl::AttestationKeySaveOutcome
SessionStoreImpl::SetWrappedAttestationKey(const SchemefulSite& site,
const Session& session,
proto::Session& session_proto,
SessionStore::SaveSessionMode mode) {
const auto& maybe_aik_id_or_error =
session.maybe_unexportable_attestation_key_id();
// The in-memory session indicates the attestation key is not yet loaded into
// the TPM by returning `ServiceError::kKeyNotReady`.
//
// During a session refresh (`kRefresh`), the refreshed session is expected
// to reuse the same attestation key. Since loading it is an expensive
// operation, we delay loading it until it is actually needed, and in the
// meantime, we preserve the existing wrapped key by copying it from the
// database entry of the old session.
//
// If this is a new session (`kNewSession`), key preservation is disabled to
// avoid leaking a key between two independent sessions.
if (mode == SessionStore::SaveSessionMode::kRefresh &&
maybe_aik_id_or_error == base::unexpected(ServiceError::kKeyNotReady)) {
proto::SiteSessions old_site_proto;
if (!session_data_->TryGetData(site.Serialize(), &old_site_proto)) {
return AttestationKeySaveOutcome::kKeyNotReadyNoSiteInDb;
}
const proto::Session* old_session =
base::FindOrNull(old_site_proto.sessions(), *session.id());
if (!old_session || !old_session->has_wrapped_attestation_key()) {
return old_session ? AttestationKeySaveOutcome::kKeyNotReadyNoOldKeyToCopy
: AttestationKeySaveOutcome::kKeyNotReadyNoSessionInDb;
}
session_proto.set_wrapped_attestation_key(
old_session->wrapped_attestation_key());
return AttestationKeySaveOutcome::kKeyNotReadyCopiedOldKey;
}
// Unexpected error (e.g. kFailure or kKeyNotFound).
ASSIGN_OR_RETURN(
std::optional<unexportable_keys::UnexportableAttestationKeyId>
maybe_aik_id,
maybe_aik_id_or_error,
[](auto) { return AttestationKeySaveOutcome::kUnexpectedError; });
// No key is expected (nullopt). Do not set it in the proto (clearing it).
if (!maybe_aik_id) {
session_proto.clear_wrapped_attestation_key();
return AttestationKeySaveOutcome::kNoAttestationKey;
}
// Wrap the attestation key and save it.
ASSIGN_OR_RETURN(std::vector<uint8_t> wrapped_attestation_key,
key_service_->GetWrappedKey(*maybe_aik_id), [](auto) {
return AttestationKeySaveOutcome::kGetWrappedKeyFailure;
});
session_proto.set_wrapped_attestation_key(
base::as_string_view(wrapped_attestation_key));
return AttestationKeySaveOutcome::kSaveSessionKeySuccess;
}
void SessionStoreImpl::DeleteSession(const SessionKey& key) {
if (db_status_ != DBStatus::kSuccess) {
return;
}
proto::SiteSessions site_proto;
std::string site_str = key.site.Serialize();
if (!session_data_->TryGetData(site_str, &site_proto)) {
return;
}
if (site_proto.sessions().count(*key.id) == 0) {
return;
}
// If this is the only session associated with the site,
// delete the site entry.
if (site_proto.mutable_sessions()->size() == 1) {
session_data_->DeleteData({site_str});
return;
}
site_proto.mutable_sessions()->erase(*key.id);
// Schedule a DB update for the site entry.
session_data_->UpdateData(key.site.Serialize(), site_proto);
}
SessionStore::SessionsMap SessionStoreImpl::GetAllSessions() const {
if (db_status_ != DBStatus::kSuccess) {
return SessionsMap();
}
// We shouldn't find invalid keys at this point, they should have all been
// filtered out in the `LoadSessions` operations. So, all session entries in
// the cache are expected to be valid.
std::vector<std::string> keys_to_delete;
std::map<std::string, proto::SiteSessions> sites_to_update;
SessionsMap all_sessions = CreateSessionsFromLoadedData(
session_data_->GetAllCached(), keys_to_delete, sites_to_update,
/*prune_expired_sessions=*/false);
CHECK(keys_to_delete.empty());
CHECK(sites_to_update.empty());
return all_sessions;
}
std::optional<proto::Session> SessionStoreImpl::GetSessionProto(
const SessionKey& session_key) const {
if (db_status_ != DBStatus::kSuccess) {
return std::nullopt;
}
proto::SiteSessions site_proto;
if (!session_data_->TryGetData(session_key.site.Serialize(), &site_proto)) {
return std::nullopt;
}
proto::Session* session =
base::FindOrNull(*site_proto.mutable_sessions(), *session_key.id);
return session ? std::optional(std::move(*session)) : std::nullopt;
}
void SessionStoreImpl::RestoreSessionBindingKey(
const SessionKey& session_key,
unexportable_keys::BackgroundTaskPriority priority,
RestoreSessionBindingKeyCallback callback) {
std::optional<proto::Session> session_proto = GetSessionProto(session_key);
session_proto ? key_service_->FromWrappedSigningKeySlowlyAsync(
base::as_byte_span(session_proto->wrapped_key()),
priority, std::move(callback))
: std::move(callback).Run(base::unexpected(
unexportable_keys::ServiceError::kKeyNotFound));
}
void SessionStoreImpl::RestoreSessionAttestationKey(
const SessionKey& session_key,
unexportable_keys::BackgroundTaskPriority priority,
RestoreSessionAttestationKeyCallback callback) {
std::optional<proto::Session> session_proto = GetSessionProto(session_key);
(session_proto && session_proto->has_wrapped_attestation_key())
? key_service_->FromWrappedAttestationKeySlowlyAsync(
base::as_byte_span(session_proto->wrapped_attestation_key()),
priority, std::move(callback))
: std::move(callback).Run(
base::unexpected(unexportable_keys::ServiceError::kKeyNotFound));
}
void SessionStoreImpl::StartGarbageCollection() {
DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
CHECK_EQ(db_status_, DBStatus::kSuccess);
key_service_->GetAllKeysForGarbageCollectionSlowlyAsync(
unexportable_keys::BackgroundTaskPriority::kBestEffort,
base::BindOnce(&SessionStoreImpl::OnGetAllKeysForGarbageCollection,
weak_ptr_factory_.GetWeakPtr()));
}
void SessionStoreImpl::OnGetAllKeysForGarbageCollection(
unexportable_keys::ServiceErrorOr<
std::vector<unexportable_keys::UnexportableSigningKeyId>>
all_key_ids_or_error) {
DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
if (!all_key_ids_or_error.has_value() || all_key_ids_or_error->empty()) {
return;
}
absl::flat_hash_set<std::vector<uint8_t>> known_wrapped_keys;
for (const auto& [_, site_sessions] : session_data_->GetAllCached()) {
for (const auto& [_, session_proto] : site_sessions.sessions()) {
if (std::string_view wrapped_key = session_proto.wrapped_key();
!wrapped_key.empty()) {
known_wrapped_keys.emplace(std::from_range, wrapped_key);
}
if (std::string_view wrapped_attestation_key =
session_proto.wrapped_attestation_key();
!wrapped_attestation_key.empty()) {
known_wrapped_keys.emplace(std::from_range, wrapped_attestation_key);
}
}
}
std::vector<unexportable_keys::UnexportableSigningKeyId> all_key_ids =
*std::move(all_key_ids_or_error);
const size_t key_count = all_key_ids.size();
base::UmaHistogramCounts100(
base::StrCat({kGarbageCollectionHistogramPrefix, "TotalKeyCount"}),
key_count);
// Don't garbage collect keys that are still used, or were created after the
// process started.
std::erase_if(
all_key_ids, [&](unexportable_keys::UnexportableSigningKeyId key_id) {
return known_wrapped_keys.contains(
key_service_->GetWrappedKey(key_id).value_or({})) ||
key_service_->GetCreationTime(key_id).value_or(
base::Time::Now()) >=
base::Process::Current().CreationTime();
});
base::UmaHistogramCounts100(
base::StrCat({kGarbageCollectionHistogramPrefix, "UsedKeyCount"}),
key_count - all_key_ids.size());
base::UmaHistogramCounts100(
base::StrCat({kGarbageCollectionHistogramPrefix, "ObsoleteKeyCount"}),
all_key_ids.size());
// Delete all remaining keys.
key_service_->DeleteKeysSlowlyAsync(
all_key_ids, unexportable_keys::BackgroundTaskPriority::kBestEffort,
base::BindOnce([](unexportable_keys::ServiceErrorOr<size_t> result) {
base::UmaHistogramCounts100(
base::StrCat({kGarbageCollectionHistogramPrefix,
"ObsoleteKeyDeletionCount"}),
result.value_or(0));
}));
}
} // namespace net::device_bound_sessions