// 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_service_impl.h"

#include <algorithm>
#include <optional>
#include <string>
#include <utility>
#include <vector>

#include "base/barrier_callback.h"
#include "base/check_deref.h"
#include "base/containers/map_util.h"
#include "base/containers/to_vector.h"
#include "base/feature_list.h"
#include "base/functional/bind.h"
#include "base/metrics/histogram_functions.h"
#include "base/process/process.h"
#include "base/task/sequenced_task_runner.h"
#include "base/time/time.h"
#include "base/types/optional_ref.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/cert/x509_certificate.h"
#include "net/cookies/cookie_partition_key_collection.h"
#include "net/cookies/cookie_store.h"
#include "net/device_bound_sessions/challenge_result.h"
#include "net/device_bound_sessions/jwk_utils.h"
#include "net/device_bound_sessions/registration_request_param.h"
#include "net/device_bound_sessions/session_binding_utils.h"
#include "net/device_bound_sessions/session_display.h"
#include "net/device_bound_sessions/session_store.h"
#include "net/ssl/ssl_cert_request_info.h"
#include "net/ssl/ssl_private_key.h"
#include "net/url_request/url_request.h"
#include "net/url_request/url_request_context.h"
#include "third_party/abseil-cpp/absl/container/flat_hash_set.h"
#include "third_party/abseil-cpp/absl/functional/overload.h"

namespace net::device_bound_sessions {

namespace {

// Parameters for the signing quota. We currently allow 6 signings in 9
// minutes per site. Reasoning:
// 1. This allows sites to refresh on average every 5 minutes, accounting for
//    proactive refreshes 2 minutes before expiry, and with some error tolerance
//    (e.g. a failed refresh or user cookie clearing) and tolerance for new
//    registration signings.
// 2. It's 6:9 instead of 2:3 to allow small bursts of login activity + new
//    registrations.
// 3. The spec notes that user agents should include quotas on registration
//    attempts to prevent identity linking for federated sessions.
constexpr size_t kSigningQuota = 6;
constexpr base::TimeDelta kSigningQuotaInterval = base::Minutes(9);

constexpr base::TimeDelta kProactiveRefreshThreshold = base::Seconds(120);

// Returns the timestamp when the session will next reach the proactive refresh
// threshold (or `base::Time::Max()` for session cookies), or `std::nullopt` if
// the session already needs a refresh.
std::optional<base::Time> GetFutureRefreshTime(base::TimeDelta lifetime) {
  if (lifetime <= kProactiveRefreshThreshold) {
    return std::nullopt;
  }
  return lifetime.is_max()
             ? base::Time::Max()
             : base::Time::Now() + (lifetime - kProactiveRefreshThreshold);
}

// Computes when the session will next reach the proactive refresh threshold.
// Returns `base::Time::Max()` for missing or session cookies because their
// expiration cannot be proactively tracked.
base::Time ComputeEarliestNextRefreshTime(base::TimeDelta lifetime) {
  return lifetime.is_zero()
             ? base::Time::Max()
             : GetFutureRefreshTime(lifetime).value_or(base::Time::Now());
}

bool SessionMatchesFilter(
    const SchemefulSite& site,
    const Session& session,
    std::optional<base::Time> created_after_time,
    std::optional<base::Time> created_before_time,
    base::RepeatingCallback<bool(const url::Origin&, const net::SchemefulSite&)>
        origin_and_site_matcher) {
  if (created_before_time && *created_before_time < session.creation_date()) {
    return false;
  }

  if (created_after_time && *created_after_time > session.creation_date()) {
    return false;
  }

  if (!origin_and_site_matcher.is_null() &&
      !origin_and_site_matcher.Run(session.origin(), site)) {
    return false;
  }

  return true;
}

class DebugHeaderBuilder {
 public:
  void AddSkippedSession(SessionKey key, RefreshResult result) {
    structured_headers::Item item;
    switch (result) {
      case RefreshResult::kRefreshed:
      // TODO(crbug.com/417401759): Add "transient_signing_error" as a supported
      // value for `Secure-Session-Skipped`.
      case RefreshResult::kTransientSigningError:
      case RefreshResult::kFatalError:
      case RefreshResult::kRefreshedAsWaiter:
      case RefreshResult::kInScopeRefreshNotYetNeeded:
        return;
      case RefreshResult::kInitializedService:
        NOTREACHED();
      case RefreshResult::kUnreachable:
        item = structured_headers::Item("unreachable",
                                        structured_headers::Item::kTokenType);
        break;
      case RefreshResult::kServerError:
        item = structured_headers::Item("server_error",
                                        structured_headers::Item::kTokenType);
        break;
      case RefreshResult::kSigningQuotaExceeded:
        item = structured_headers::Item("quota_exceeded",
                                        structured_headers::Item::kTokenType);
        break;
    }

    structured_headers::Parameters params = {
        {"session_identifier", structured_headers::Item(key.id.value())}};
    skipped_sessions_.emplace_back(std::move(item), std::move(params));
  }

  std::optional<std::string> Build() {
    if (skipped_sessions_.empty()) {
      return std::nullopt;
    }

    return structured_headers::SerializeList(std::move(skipped_sessions_));
  }

 private:
  structured_headers::List skipped_sessions_;
};

bool IsProactiveRefreshCandidate(
    Session& existing_session,
    const Session& new_session,
    const CookieAndLineAccessResultList& maybe_stored_cookies) {
  // Get the shortest lifetime of a bound cookie set by the current
  // refresh request. This assumes:
  // 1. The current refresh sets all bound cookies
  // 2. The proactive refresh would have set the same lifetimes
  // These assumptions are good enough for histogram logging, but likely
  // not true for all sites.
  base::Time current_time = base::Time::Now();
  base::TimeDelta minimum_lifetime = base::TimeDelta::Max();
  for (const CookieCraving& cookie_craving : new_session.cookies()) {
    for (const CookieAndLineWithAccessResult& cookie_and_line :
         maybe_stored_cookies) {
      if (cookie_and_line.cookie.has_value() &&
          cookie_craving.IsSatisfiedBy(cookie_and_line.cookie.value())) {
        minimum_lifetime =
            std::min(minimum_lifetime,
                     cookie_and_line.cookie->ExpiryDate() - current_time);
      }
    }
  }

  base::UmaHistogramLongTimes100(
      "Net.DeviceBoundSessions.MinimumBoundCookieLifetime", minimum_lifetime);

  std::optional<base::Time> last_proactive_refresh_opportunity =
      existing_session.TakeLastProactiveRefreshOpportunity();

  if (!last_proactive_refresh_opportunity.has_value()) {
    return false;
  }

  return minimum_lifetime >= current_time - *last_proactive_refresh_opportunity;
}

void LogProactiveRefreshAttempt(
    SessionServiceImpl::ProactiveRefreshAttempt attempt) {
  base::UmaHistogramEnumeration(
      "Net.DeviceBoundSessions.ProactiveRefreshAttempt", attempt);
}

}  // namespace

SessionServiceImpl::SessionServiceImpl(
    unexportable_keys::UnexportableKeyService& key_service,
    const URLRequestContext* request_context,
    SessionStore* store,
    const std::vector<SchemefulSite>& restricted_sites,
    CookieAccessCallback has_cookie_access_cb,
    SelectClientCertificateHandler client_cert_handler)
    : pending_initialization_(!!store),
      key_service_(key_service),
      context_(request_context),
      session_store_(store),
      restricted_sites_(restricted_sites),
      client_cert_handler_(std::move(client_cert_handler)),
      has_cookie_access_cb_(std::move(has_cookie_access_cb)) {
  ignore_signing_quota_ = !features::kDeviceBoundSessionsSigningQuota.Get();
  CHECK(context_);
  CHECK(client_cert_handler_);
}

SessionServiceImpl::~SessionServiceImpl() = default;

void SessionServiceImpl::SelectClientCertificate(
    const GURL& url,
    scoped_refptr<SSLCertRequestInfo> cert_info,
    SelectClientCertificateCallback callback) {
  client_cert_handler_.Run(url, std::move(cert_info), std::move(callback));
}

void SessionServiceImpl::PrewarmSessionsForUrl(const GURL& url,
                                               PrewarmCallback callback) {
  if (pending_initialization_) {
    queued_operations_.push_back(
        base::BindOnce(&SessionServiceImpl::PrewarmSessionsForUrl,
                       weak_factory_.GetWeakPtr(), url, std::move(callback)));
    return;
  }

  if (!url.is_valid() || !context_ || !context_->cookie_store()) {
    base::SequencedTaskRunner::GetCurrentDefault()->PostTask(
        FROM_HERE, base::BindOnce(std::move(callback), SessionPrewarmResult{}));
    return;
  }

  SchemefulSite site(url);
  std::vector<SessionKey> matching_sessions;
  for (const auto& [session_key, session] : GetSessionsForSite(site)) {
    if (session->IncludesUrl(url)) {
      matching_sessions.push_back(session_key);
    }
  }

  if (matching_sessions.empty()) {
    base::SequencedTaskRunner::GetCurrentDefault()->PostTask(
        FROM_HERE, base::BindOnce(std::move(callback), SessionPrewarmResult{}));
    return;
  }

  context_->cookie_store()->GetCookieListWithOptionsAsync(
      url, CookieOptions::MakeAllInclusive(), CookiePartitionKeyCollection(),
      base::BindOnce(&SessionServiceImpl::OnGetCookiesForPrewarm,
                     weak_factory_.GetWeakPtr(), url,
                     std::move(matching_sessions), std::move(callback)));
}

void SessionServiceImpl::OnGetCookiesForPrewarm(
    const GURL& url,
    std::vector<SessionKey> matching_sessions,
    PrewarmCallback callback,
    const CookieAccessResultList& cookies,
    const CookieAccessResultList& excluded_cookies) {
  auto barrier = base::BarrierCallback<PrewarmResult>(
      matching_sessions.size(),
      base::BindOnce(&SessionServiceImpl::OnAllPrewarmSessionsDone,
                     weak_factory_.GetWeakPtr(), std::move(callback)));

  for (const SessionKey& session_key : matching_sessions) {
    Session* session = GetSession(session_key);
    if (!session) {
      barrier.Run({.result = RefreshResult::kFatalError});
      continue;
    }

    NotifySessionAccess(base::NullCallback(),
                        SessionAccess::AccessType::kUpdate, session_key,
                        *session);

    if (session->ShouldBackoff()) {
      barrier.Run({.result = RefreshResult::kUnreachable});
      continue;
    }

    if (SigningQuotaExceeded(session_key.site)) {
      barrier.Run({.result = RefreshResult::kSigningQuotaExceeded});
      continue;
    }

    if (std::optional<base::Time> refresh_time = GetFutureRefreshTime(
            session->MinimumBoundCookieLifetime(cookies))) {
      barrier.Run({.result = RefreshResult::kInScopeRefreshNotYetNeeded,
                   .earliest_next_refresh_time = *refresh_time});
      continue;
    }

    // Refresh needed (either missing cookie or expiring cookie below
    // threshold).
    auto [it, inserted] = proactive_requests_.try_emplace(session_key);
    it->second.completion_callbacks.push_back(barrier);
    if (inserted && !deferred_requests_.contains(session_key)) {
      url::Origin origin = url::Origin::Create(url);
      StartSessionRefresh(
          session_key,
          {
              .trigger = RefreshTrigger::kProactive,
              .isolation_info =
                  net::IsolationInfo::CreateForInternalRequest(origin),
              .site_for_cookies = net::SiteForCookies::FromOrigin(origin),
              .initiator = origin,
              .priority =
                  unexportable_keys::BackgroundTaskPriority::kBestEffort,
          });
    }
  }
}

void SessionServiceImpl::CompleteProactiveRefresh(
    const SessionKey& session_key,
    RefreshResult result,
    std::vector<base::OnceCallback<void(PrewarmResult)>> callbacks) {
  if (callbacks.empty()) {
    return;
  }

  Session* session = GetSession(session_key);
  if (result == RefreshResult::kRefreshed && session &&
      context_->cookie_store()) {
    context_->cookie_store()->GetCookieListWithOptionsAsync(
        session->refresh_url(), CookieOptions::MakeAllInclusive(),
        CookiePartitionKeyCollection(),
        base::BindOnce(&SessionServiceImpl::OnGetCookiesAfterProactiveRefresh,
                       weak_factory_.GetWeakPtr(), session_key,
                       std::move(callbacks), result));
    return;
  }

  PrewarmResult prewarm_result{
      .result = (!session && result == RefreshResult::kRefreshed)
                    ? RefreshResult::kFatalError
                    : result};
  for (auto& callback : callbacks) {
    std::move(callback).Run(prewarm_result);
  }
}

void SessionServiceImpl::OnGetCookiesAfterProactiveRefresh(
    const SessionKey& session_key,
    std::vector<base::OnceCallback<void(PrewarmResult)>> callbacks,
    RefreshResult refresh_result,
    const CookieAccessResultList& cookies,
    const CookieAccessResultList& excluded_cookies) {
  Session* session = GetSession(session_key);
  PrewarmResult result =
      session
          ? PrewarmResult{.result = refresh_result,
                          .earliest_next_refresh_time =
                              ComputeEarliestNextRefreshTime(
                                  session->MinimumBoundCookieLifetime(cookies))}
          : PrewarmResult{.result = RefreshResult::kFatalError};

  for (auto& callback : callbacks) {
    std::move(callback).Run(result);
  }
}

void SessionServiceImpl::OnAllPrewarmSessionsDone(
    PrewarmCallback callback,
    std::vector<PrewarmResult> session_results) {
  SessionPrewarmResult final_result;
  final_result.results.reserve(session_results.size());
  for (const auto& sr : session_results) {
    final_result.results.push_back(sr.result);
    final_result.earliest_next_refresh_time = std::min(
        final_result.earliest_next_refresh_time, sr.earliest_next_refresh_time);
  }

  std::move(callback).Run(std::move(final_result));
}

void SessionServiceImpl::LoadSessionsAsync() {
  if (!session_store_) {
    return;
  }
  session_store_->LoadSessions(base::BindOnce(
      &SessionServiceImpl::OnLoadSessionsComplete, weak_factory_.GetWeakPtr()));
}

void SessionServiceImpl::RegisterBoundSession(
    OnAccessCallback on_access_callback,
    RegistrationFetcherParam registration_params,
    const IsolationInfo& isolation_info,
    const net::SiteForCookies& site_for_cookies,
    const NetLogWithSource& net_log,
    const std::optional<url::Origin>& original_request_initiator) {
  if (const auto& provider_params = registration_params.provider_params();
      provider_params.has_value() &&
      provider_params->provider_session_id.has_value()) {
    if (!base::FeatureList::IsEnabled(
            features::kDeviceBoundSessionsFederatedRegistration)) {
      // Simply ignore headers with a provider_session_id if the flag
      // isn't enabled.
      return;
    }

    // Copy provider params before `registration_params` gets `std::move()`d.
    ProviderRegistrationParams params = *provider_params;
    GetFederatedProviderSessionIfValid(
        std::move(params), on_access_callback,
        base::BindOnce(&SessionServiceImpl::RegisterBoundSessionInternal,
                       weak_factory_.GetWeakPtr(), on_access_callback,
                       std::move(registration_params), isolation_info,
                       site_for_cookies, net_log, original_request_initiator));
    return;
  }

  RegisterBoundSessionInternal(
      std::move(on_access_callback), std::move(registration_params),
      isolation_info, site_for_cookies, net_log, original_request_initiator,
      /*federated_provider_session=*/nullptr);
}

void SessionServiceImpl::RegisterBoundSessionInternal(
    OnAccessCallback on_access_callback,
    RegistrationFetcherParam registration_params,
    const IsolationInfo& isolation_info,
    const net::SiteForCookies& site_for_cookies,
    const NetLogWithSource& net_log,
    const std::optional<url::Origin>& original_request_initiator,
    base::expected<Session*, SessionError> federated_provider_session) {
  bool is_google_subdomain_for_histograms = IsSubdomainOf(
      registration_params.registration_endpoint().host(), "google.com");
  SchemefulSite site =
      SchemefulSite(registration_params.registration_endpoint());
  // A federated session was attempted but had an error.
  if (!federated_provider_session.has_value()) {
    OnRegistrationComplete(
        std::move(on_access_callback), is_google_subdomain_for_histograms,
        /*is_federated_registration_for_histograms=*/true, site,
        /*fetcher=*/nullptr,
        RegistrationResult(std::move(federated_provider_session.error())));
    return;
  }

  if (*federated_provider_session) {
    Session* provider_session = *federated_provider_session;
    SessionKey provider_session_key{SchemefulSite(provider_session->origin()),
                                    provider_session->id()};
    NotifySessionAccess(on_access_callback, SessionAccess::AccessType::kUpdate,
                        provider_session_key, *provider_session);
  }

  net::NetLogSource net_log_source_for_registration = net::NetLogSource(
      net::NetLogSourceType::URL_REQUEST, net::NetLog::Get()->NextID());
  net_log.AddEventReferencingSource(
      net::NetLogEventType::DBSC_REGISTRATION_REQUEST,
      net_log_source_for_registration);

  std::vector<crypto::SignatureVerifier::SignatureAlgorithm> supported_algos =
      base::ToVector(registration_params.supported_algos());
  RegistrationRequestParam request_params =
      RegistrationRequestParam::CreateForRegistration(
          std::move(registration_params));
  std::unique_ptr<RegistrationFetcher> fetcher =
      RegistrationFetcher::CreateFetcher(
          request_params, *this, key_service_.get(), context_.get(),
          isolation_info, site_for_cookies, net_log_source_for_registration,
          original_request_initiator,
          unexportable_keys::BackgroundTaskPriority::kBestEffort);
  RegistrationFetcher* fetcher_raw = fetcher.get();
  registration_fetchers_.insert(std::move(fetcher));

  auto callback = base::BindOnce(
      &SessionServiceImpl::OnRegistrationComplete, weak_factory_.GetWeakPtr(),
      std::move(on_access_callback), is_google_subdomain_for_histograms,
      /*is_federated_registration_for_histograms=*/federated_provider_session !=
          nullptr,
      site);
  if (*federated_provider_session) {
    Session* provider_session = *federated_provider_session;
    fetcher_raw->StartFetchWithFederatedKey(
        request_params, *provider_session->unexportable_key_id(),
        provider_session->origin().GetURL(), std::move(callback));
    // `fetcher_raw` may be deleted.
  } else {
    fetcher_raw->StartCreateTokenAndFetch(request_params, supported_algos,
                                          std::move(callback));
    // `fetcher_raw` may be deleted.
  }
}

void SessionServiceImpl::GetFederatedProviderSessionIfValid(
    ProviderRegistrationParams provider_params,
    OnAccessCallback on_access_callback,
    base::OnceCallback<void(base::expected<Session*, SessionError>)> callback) {
  CHECK(provider_params.provider_session_id.has_value());
  const GURL& provider_url = provider_params.provider_url;
  // This is a federated session registration.
  if (!provider_url.is_valid() || url::Origin::Create(provider_url).opaque()) {
    std::move(callback).Run(base::unexpected(
        SessionError(SessionError::kInvalidFederatedSessionUrl)));
    return;
  }

  SessionKey provider_session_key{SchemefulSite(provider_url),
                                  *provider_params.provider_session_id};
  Session* provider_session = GetSession(provider_session_key);
  if (!provider_session) {
    std::move(callback).Run(base::unexpected(SessionError(
        SessionError::kInvalidFederatedSessionProviderSessionMissing)));
    return;
  }

  if (url::Origin::Create(provider_url) != provider_session->origin()) {
    std::move(callback).Run(base::unexpected(SessionError(
        SessionError::kInvalidFederatedSessionWrongProviderOrigin)));
    return;
  }

  if (!provider_session->unexportable_key_id().has_value()) {
    RestoreSessionKey(
        provider_session_key, on_access_callback,
        base::BindOnce(&SessionServiceImpl::CheckFederatedProviderKey,
                       weak_factory_.GetWeakPtr(), provider_session_key,
                       std::move(provider_params.provider_key),
                       std::move(callback)));
    return;
  }

  CheckFederatedProviderKey(
      std::move(provider_session_key), std::move(provider_params.provider_key),
      std::move(callback), *provider_session->unexportable_key_id());
}

void SessionServiceImpl::CheckFederatedProviderKey(
    SessionKey provider_session_key,
    std::string provider_key_thumbprint,
    base::OnceCallback<void(base::expected<Session*, SessionError>)> callback,
    std::optional<unexportable_keys::UnexportableSigningKeyId> provider_key) {
  if (!provider_key) {
    // Failed to restore provider key.
    std::move(callback).Run(base::unexpected(SessionError(
        SessionError::kInvalidFederatedSessionProviderFailedToRestoreKey)));
    return;
  }

  Session* provider_session = GetSession(provider_session_key);
  if (!provider_session) {
    // Provider session not found, fail the registration.
    std::move(callback).Run(base::unexpected(SessionError(
        SessionError::kInvalidFederatedSessionProviderSessionMissing)));
    return;
  }

  unexportable_keys::ServiceErrorOr<
      crypto::SignatureVerifier::SignatureAlgorithm>
      algorithm =
          key_service_->GetAlgorithm(*provider_session->unexportable_key_id());
  if (!algorithm.has_value()) {
    std::move(callback).Run(
        base::unexpected(SessionError(SessionError::kInvalidFederatedKey)));
    return;
  }

  unexportable_keys::ServiceErrorOr<std::vector<uint8_t>> pub_key =
      key_service_->GetSubjectPublicKeyInfo(
          *provider_session->unexportable_key_id());
  if (!pub_key.has_value()) {
    std::move(callback).Run(
        base::unexpected(SessionError(SessionError::kInvalidFederatedKey)));
    return;
  }

  std::string thumbprint = CreateJwkThumbprint(*algorithm, *pub_key);
  if (thumbprint != provider_key_thumbprint) {
    std::move(callback).Run(base::unexpected(
        SessionError(SessionError::kFederatedKeyThumbprintMismatch)));
    return;
  }

  std::move(callback).Run(provider_session);
}

void SessionServiceImpl::OnLoadSessionsComplete(
    SessionStore::SessionsMap sessions) {
  unpartitioned_sessions_.merge(sessions);
  pending_initialization_ = false;

  std::vector<base::OnceClosure> queued_operations =
      std::move(queued_operations_);
  for (base::OnceClosure& closure : queued_operations) {
    std::move(closure).Run();
  }

  base::UmaHistogramCounts1000(
      "Net.DeviceBoundSessions.RequestsDeferredForInitialization",
      requests_before_initialization_);
}

void SessionServiceImpl::OnRegistrationComplete(
    OnAccessCallback on_access_callback,
    bool is_google_subdomain_for_histograms,
    bool is_federated_registration_for_histograms,
    SchemefulSite site,
    RegistrationFetcher* fetcher,
    RegistrationResult registration_result) {
  if (is_google_subdomain_for_histograms) {
    base::UmaHistogramBoolean(
        "Net.DeviceBoundSessions.GoogleRegistrationIsFromStandard", true);
  }
  SessionError::ErrorType result =
      OnRegistrationCompleteInternal(std::move(on_access_callback), fetcher,
                                     std::move(registration_result), site);
  base::UmaHistogramEnumeration("Net.DeviceBoundSessions.RegistrationResult",
                                result);
  if (is_federated_registration_for_histograms) {
    base::UmaHistogramEnumeration(
        "Net.DeviceBoundSessions.RegistrationResult.Federated", result);
  } else {
    base::UmaHistogramEnumeration(
        "Net.DeviceBoundSessions.RegistrationResult.Standalone", result);
  }
}

std::ranges::subrange<SessionServiceImpl::SessionsMap::iterator>
SessionServiceImpl::GetSessionsForSite(const SchemefulSite& site) {
  const auto now = base::Time::Now();
  // Session keys are sorted by site, then identifier. So the first
  // element not less than (`site`, "") is the first session for this
  // site.
  auto it =
      unpartitioned_sessions_.lower_bound(SessionKey{site, Session::Id("")});
  while (it != unpartitioned_sessions_.end() && it->first.site == site) {
    auto curit = it;
    ++it;

    const auto& [session_key, session] = *curit;
    if (now >= session->expiry_date()) {
      // Since this deletion is not due to a request, we do not need to
      // provide a per-request callback here.
      DeleteSessionAndNotifyInternal(DeletionReason::kExpired, curit,
                                     base::NullCallback());
    } else {
      session->RecordAccess();
    }
  }

  return std::ranges::subrange<SessionsMap::iterator>(
      unpartitioned_sessions_.lower_bound(SessionKey{site, Session::Id("")}),
      it);
}

std::optional<SessionService::DeferralParams> SessionServiceImpl::ShouldDefer(
    DbscRequest& request,
    HttpRequestHeaders* extra_headers,
    const FirstPartySetMetadata& first_party_set_metadata) {
  if (request.device_bound_session_mode() ==
          net::DeviceBoundSessionMode::kDisabled ||
      request.device_bound_session_mode() ==
          net::DeviceBoundSessionMode::kBypassDeferral) {
    return std::nullopt;
  }

  if (pending_initialization_) {
    return DeferralParams();
  }

  SchemefulSite site(request.url());
  DebugHeaderBuilder debug_header_builder;
  const base::flat_map<SessionKey, RefreshResult>& previous_deferrals =
      request.device_bound_session_deferrals();
  for (const auto& [session_key, session] : GetSessionsForSite(site)) {
    MaybeIncreaseSessionUsage(session_key, request,
                              SessionUsage::kNoSiteMatchNotInScope);

    if (!session->IsInScope(request)) {
      continue;
    }

    base::TimeDelta minimum_lifetime = session->MinimumBoundCookieLifetime(
        request, first_party_set_metadata, session_key);
    if (minimum_lifetime.is_zero()) {
      auto previous_deferrals_it = previous_deferrals.find(session_key);
      if (previous_deferrals_it != previous_deferrals.end() &&
          previous_deferrals_it->second != RefreshResult::kRefreshedAsWaiter) {
        debug_header_builder.AddSkippedSession(previous_deferrals_it->first,
                                               previous_deferrals_it->second);
        continue;
      }

      NotifySessionAccess(request.device_bound_session_access_callback(),
                          SessionAccess::AccessType::kUpdate, session_key,
                          *session);
      return DeferralParams(session->id());
    }

    MaybeStartProactiveRefresh(request.device_bound_session_access_callback(),
                               request, session_key, minimum_lifetime);
  }

  std::optional<std::string> debug_header = debug_header_builder.Build();
  if (debug_header.has_value()) {
    extra_headers->SetHeader("Secure-Session-Skipped", *debug_header);
  }

  return std::nullopt;
}

void SessionServiceImpl::DeferRequestForRefresh(
    DbscRequest& request,
    DeferralParams deferral,
    RefreshCompleteCallback callback) {
  CHECK(callback);

  if (deferral.is_pending_initialization) {
    CHECK(pending_initialization_);
    requests_before_initialization_++;
    // Due to the need to recompute `first_party_set_metadata`, we always
    // restart the request after initialization completes.
    queued_operations_.push_back(base::BindOnce(
        std::move(callback), RefreshResult::kInitializedService));
    return;
  }

  SessionKey session_key{SchemefulSite(request.url()), *deferral.session_id};
  // For the first deferring request, create a new vector and add the request.
  auto [it, inserted] = deferred_requests_.try_emplace(session_key);
  // Add the request callback to the deferred list.
  it->second.push_back(
      {.request = request.GetWeakPtr(), .callback = std::move(callback)});

  auto* session = GetSession(session_key);
  CHECK(session);
  // Notify the request that it has been deferred for refreshed cookies.
  NotifySessionAccess(request.device_bound_session_access_callback(),
                      SessionAccess::AccessType::kUpdate, session_key,
                      *session);
  bool deferred_request_refresh_in_progress = !inserted;
  base::UmaHistogramBoolean(
      "Net.DeviceBoundSessions."
      "DeferredRequestRefreshAlreadyInProgressOnDeferAttempt",
      deferred_request_refresh_in_progress);
  if (deferred_request_refresh_in_progress) {
    return;
  }
  bool proactive_refresh_in_progress =
      proactive_requests_.find(session_key) != proactive_requests_.end();
  base::UmaHistogramBoolean(
      "Net.DeviceBoundSessions.ProactiveRefreshAlreadyInProgressOnDeferAttempt",
      proactive_refresh_in_progress);
  if (proactive_refresh_in_progress) {
    return;
  }

  if (session->ShouldBackoff()) {
    UnblockWaitingRequests(session_key, RefreshResult::kUnreachable);
    return;
  }

  it->second.back().triggered_refresh = true;
  StartSessionRefresh(
      session_key,
      {
          .trigger = RefreshTrigger::kMissingCookie,
          .isolation_info = request.isolation_info(),
          .site_for_cookies = request.site_for_cookies(),
          .initiator = request.initiator(),
          .priority = unexportable_keys::BackgroundTaskPriority::kUserBlocking,
          .access_callback = request.device_bound_session_access_callback(),
          .net_log = request.net_log(),
      });
}

void SessionServiceImpl::OnRefreshRequestCompletion(
    RefreshTrigger trigger,
    OnAccessCallback on_access_callback,
    SessionKey session_key,
    RegistrationFetcher* fetcher,
    RegistrationResult registration_result) {
  SessionError::ErrorType result = OnRefreshRequestCompletionInternal(
      std::move(on_access_callback), session_key, fetcher,
      std::move(registration_result));

  Session* session = GetSession(session_key);
  if (session) {
    session->InformOfRefreshResult(
        /*was_proactive=*/trigger != RefreshTrigger::kMissingCookie, result);
  }

  std::string histogram_base = "Net.DeviceBoundSessions.RefreshResult";
  std::string suffix;
  switch (trigger) {
    case RefreshTrigger::kProactive:
      suffix = ".Proactive";
      break;
    case RefreshTrigger::kMissingCookie:
      suffix = ".MissingCookie";
      break;
  }
  base::UmaHistogramEnumeration(histogram_base, result);
  base::UmaHistogramEnumeration(histogram_base + suffix, result);
}

// Continue or restart all deferred requests and complete any proactive
// refresh requests waiting for the session, removing the session key from
// both maps.
void SessionServiceImpl::UnblockWaitingRequests(
    const SessionKey& session_key,
    RefreshResult result,
    std::optional<net::device_bound_sessions::SessionError> fetch_error,
    std::optional<SessionDisplay> new_session_display,
    std::optional<bool> is_proactive_refresh_candidate,
    std::optional<base::TimeDelta> minimum_proactive_refresh_threshold) {
  bool has_proactive_request = false;
  if (auto node = proactive_requests_.extract(session_key)) {
    has_proactive_request = true;
    base::UmaHistogramTimes("Net.DeviceBoundSessions.ProactiveRefreshDuration",
                            node.mapped().timer.Elapsed());
    CompleteProactiveRefresh(session_key, result,
                             std::move(node.mapped().completion_callbacks));
  }

  auto it = deferred_requests_.find(session_key);
  bool has_deferred_request = it != deferred_requests_.end();

  if (has_proactive_request || has_deferred_request) {
    NotifyIfEventCallbackListeners([&] {
      return SessionEvent::MakeRefreshEvent(
          session_key.site, session_key.id.value(),
          /*succeeded=*/result == RefreshResult::kRefreshed, result,
          std::move(fetch_error), std::move(new_session_display),
          /*was_fully_proactive_refresh=*/!has_deferred_request);
    });
  }

  if (!has_deferred_request) {
    return;
  }

  auto requests = std::move(it->second);
  deferred_requests_.erase(it);

  base::UmaHistogramCounts100("Net.DeviceBoundSessions.RequestDeferredCount",
                              requests.size());

  if (is_proactive_refresh_candidate.has_value() &&
      minimum_proactive_refresh_threshold.has_value()) {
    base::UmaHistogramLongTimes100(
        "Net.DeviceBoundSessions.MinimumProactiveRefreshThreshold",
        *minimum_proactive_refresh_threshold);
    if (*is_proactive_refresh_candidate) {
      base::UmaHistogramLongTimes100(
          "Net.DeviceBoundSessions.MinimumProactiveRefreshThreshold.Success",
          *minimum_proactive_refresh_threshold);
    } else {
      base::UmaHistogramLongTimes100(
          "Net.DeviceBoundSessions.MinimumProactiveRefreshThreshold.Failure",
          *minimum_proactive_refresh_threshold);
    }

    if (*is_proactive_refresh_candidate) {
      if (*minimum_proactive_refresh_threshold <= base::Seconds(30)) {
        base::UmaHistogramCounts100(
            "Net.DeviceBoundSessions.ProactiveRefreshCandidateDeferredCount."
            "ThirtySeconds",
            requests.size());
        for (auto& request : requests) {
          base::UmaHistogramTimes(
              "Net.DeviceBoundSessions."
              "ProactiveRefreshCandidateRequestDeferredDuration.ThirtySeconds",
              request.timer.Elapsed());
        }
      }

      if (*minimum_proactive_refresh_threshold <= base::Minutes(1)) {
        base::UmaHistogramCounts100(
            "Net.DeviceBoundSessions.ProactiveRefreshCandidateDeferredCount."
            "OneMinute",
            requests.size());
        for (auto& request : requests) {
          base::UmaHistogramTimes(
              "Net.DeviceBoundSessions."
              "ProactiveRefreshCandidateRequestDeferredDuration.OneMinute",
              request.timer.Elapsed());
        }
      }

      if (*minimum_proactive_refresh_threshold <= base::Minutes(2)) {
        base::UmaHistogramCounts100(
            "Net.DeviceBoundSessions.ProactiveRefreshCandidateDeferredCount."
            "TwoMinutes",
            requests.size());
        for (auto& request : requests) {
          base::UmaHistogramTimes(
              "Net.DeviceBoundSessions."
              "ProactiveRefreshCandidateRequestDeferredDuration.TwoMinutes",
              request.timer.Elapsed());
        }
      }
    }
  }

  for (auto& request : requests) {
    base::UmaHistogramTimes("Net.DeviceBoundSessions.RequestDeferredDuration",
                            request.timer.Elapsed());
    base::UmaHistogramEnumeration("Net.DeviceBoundSessions.DeferralResult",
                                  result);
    if (request.timer.Elapsed() <= base::Milliseconds(1)) {
      base::UmaHistogramEnumeration(
          "Net.DeviceBoundSessions.DeferralResult.Instant", result);
    } else {
      base::UmaHistogramEnumeration(
          "Net.DeviceBoundSessions.DeferralResult.Slow", result);
    }
    RefreshResult final_result = result;
    if (result == RefreshResult::kRefreshed && !request.triggered_refresh) {
      final_result = RefreshResult::kRefreshedAsWaiter;
    }
    std::move(request.callback).Run(final_result);
  }
}

void SessionServiceImpl::SetChallengeForBoundSession(
    OnAccessCallback on_access_callback,
    DbscRequest& request,
    const FirstPartySetMetadata& first_party_set_metadata,
    const SessionChallengeParam& param) {
  ChallengeResult result = SetChallengeForBoundSessionInternal(
      std::move(on_access_callback), request, first_party_set_metadata, param);
  NotifyIfEventCallbackListeners([&] {
    return SessionEvent::MakeChallengeEvent(
        SchemefulSite(request.url()), param.session_id(),
        result == ChallengeResult::kSuccess, result, param.challenge());
  });
}

ChallengeResult SessionServiceImpl::SetChallengeForBoundSessionInternal(
    OnAccessCallback on_access_callback,
    DbscRequest& request,
    const FirstPartySetMetadata& first_party_set_metadata,
    const SessionChallengeParam& param) {
  if (!param.session_id()) {
    return ChallengeResult::kNoSessionId;
  }

  SessionKey session_key{SchemefulSite(request.url()),
                         Session::Id(*param.session_id())};
  Session* session = GetSession(session_key);
  if (!session) {
    return ChallengeResult::kNoSessionMatch;
  }

  if (!session->CanSetBoundCookie(request, first_party_set_metadata)) {
    return ChallengeResult::kCantSetBoundCookie;
  }

  NotifySessionAccess(on_access_callback, SessionAccess::AccessType::kUpdate,
                      session_key, *session);
  session->set_cached_challenge(param.challenge());
  return ChallengeResult::kSuccess;
}

void SessionServiceImpl::NotifyIfEventCallbackListeners(
    base::FunctionRef<SessionEvent()> event_creator) {
  if (!event_callbacks_.empty()) {
    event_callbacks_.Notify(event_creator());
  }
}

void SessionServiceImpl::GetAllSessionsAsync(
    base::OnceCallback<void(const std::vector<SessionKey>&)> callback) {
  if (pending_initialization_) {
    queued_operations_.push_back(base::BindOnce(
        &SessionServiceImpl::GetAllSessionsAsync,
        // `base::Unretained` is safe because the callback is stored in
        // `queued_operations_`, which is owned by `this`.
        base::Unretained(this), std::move(callback)));
  } else {
    std::vector<SessionKey> sessions = base::ToVector(
        unpartitioned_sessions_, [](const auto& pair) { return pair.first; });
    base::SequencedTaskRunner::GetCurrentDefault()->PostTask(
        FROM_HERE, base::BindOnce(std::move(callback), std::move(sessions)));
  }
}

void SessionServiceImpl::GetAllSessionDisplaysAsync(
    base::OnceCallback<void(const std::vector<SessionDisplay>&)> callback) {
  if (pending_initialization_) {
    queued_operations_.push_back(base::BindOnce(
        &SessionServiceImpl::GetAllSessionDisplaysAsync,
        // `base::Unretained` is safe because the callback is stored in
        // `queued_operations_`, which is owned by `this`.
        base::Unretained(this), std::move(callback)));
  } else {
    std::vector<SessionDisplay> session_displays = base::ToVector(
        unpartitioned_sessions_,
        [](const auto& pair) { return pair.second->ToDisplay(); });
    base::SequencedTaskRunner::GetCurrentDefault()->PostTask(
        FROM_HERE,
        base::BindOnce(std::move(callback), std::move(session_displays)));
  }
}

void SessionServiceImpl::DeleteSessionAndNotify(
    DeletionReason reason,
    const SessionKey& session_key,
    SessionService::OnAccessCallback per_request_callback) {
  if (pending_initialization_) {
    queued_operations_.push_back(base::BindOnce(
        &SessionServiceImpl::DeleteSessionAndNotify,
        // `base::Unretained` is safe because the callback is stored in
        // `queued_operations_`, which is owned by `this`.
        base::Unretained(this), reason, session_key,
        std::move(per_request_callback)));
    return;
  }

  auto it = unpartitioned_sessions_.find(session_key);
  if (it == unpartitioned_sessions_.end()) {
    return;
  }

  DeleteSessionAndNotifyInternal(reason, it, per_request_callback);
}

const Session* SessionServiceImpl::GetSession(
    const SessionKey& session_key) const {
  auto it = unpartitioned_sessions_.find(session_key);
  if (it != unpartitioned_sessions_.end()) {
    return it->second.get();
  }
  return nullptr;
}

Session* SessionServiceImpl::GetSession(const SessionKey& session_key) {
  return const_cast<Session*>(std::as_const(*this).GetSession(session_key));
}

void SessionServiceImpl::AddSession(
    const SchemefulSite& site,
    SessionParams params,
    base::span<const uint8_t> wrapped_key,
    base::OnceCallback<void(SessionError::ErrorType)> callback) {
  key_service_->FromWrappedSigningKeySlowlyAsync(
      wrapped_key, unexportable_keys::BackgroundTaskPriority::kBestEffort,
      base::BindOnce(&SessionServiceImpl::OnAddSessionKeyRestored,
                     weak_factory_.GetWeakPtr(), site, std::move(params),
                     std::move(callback)));
}

const SessionService::SignedRefreshChallenge*
SessionServiceImpl::GetLatestSignedRefreshChallenge(
    const SessionKey& session_key) {
  auto signed_challenge_it =
      latest_signed_refresh_challenges_.find(session_key);
  if (signed_challenge_it == latest_signed_refresh_challenges_.end()) {
    return nullptr;
  }
  return &signed_challenge_it->second;
}

void SessionServiceImpl::SetLatestSignedRefreshChallenge(
    SessionKey session_key,
    SessionService::SignedRefreshChallenge signed_refresh_challenge) {
  latest_signed_refresh_challenges_[std::move(session_key)] =
      std::move(signed_refresh_challenge);
}

base::expected<std::unique_ptr<Session>, SessionError::ErrorType>
SessionServiceImpl::CreateSessionFromUnexportableKey(
    SessionParams params,
    unexportable_keys::ServiceErrorOr<
        unexportable_keys::UnexportableSigningKeyId> key_or_error) {
  if (!key_or_error.has_value()) {
    return base::unexpected(SessionError::kFailedToUnwrapKey);
  }
  params.key_id = *key_or_error;
  base::expected<std::unique_ptr<net::device_bound_sessions::Session>,
                 net::device_bound_sessions::SessionError>
      session_or_error =
          net::device_bound_sessions::Session::CreateIfValid(params);
  if (!session_or_error.has_value()) {
    return base::unexpected(session_or_error.error().type);
  }
  return std::move(session_or_error.value());
}

void SessionServiceImpl::OnAddSessionKeyRestored(
    const SchemefulSite& site,
    SessionParams params,
    base::OnceCallback<void(SessionError::ErrorType)> callback,
    unexportable_keys::ServiceErrorOr<
        unexportable_keys::UnexportableSigningKeyId> key_or_error) {
  base::expected<std::unique_ptr<Session>, SessionError::ErrorType>
      session_or_error = CreateSessionFromUnexportableKey(
          std::move(params), std::move(key_or_error));

  NotifyIfEventCallbackListeners([&] {
    bool succeeded = session_or_error.has_value();
    SessionError::ErrorType result =
        succeeded ? SessionError::kSuccess : session_or_error.error();
    std::optional<std::string> session_id;
    std::optional<SessionDisplay> display_info;
    if (succeeded) {
      session_id = session_or_error.value()->id().value();
      display_info = session_or_error.value()->ToDisplay();
    }
    return SessionEvent::MakeCreationEvent(site, std::move(session_id),
                                           succeeded, SessionError(result),
                                           std::move(display_info));
  });

  if (!session_or_error.has_value()) {
    std::move(callback).Run(session_or_error.error());
    return;
  }

  NotifySessionAccess(base::NullCallback(),
                      SessionAccess::AccessType::kCreation,
                      SessionKey{site, session_or_error.value()->id()},
                      *session_or_error.value());

  AddSession(site, std::move(session_or_error.value()));
  std::move(callback).Run(SessionError::kSuccess);
}

void SessionServiceImpl::AddSession(const SchemefulSite& site,
                                    std::unique_ptr<Session> session,
                                    SessionStore::SaveSessionMode mode) {
  if (session_store_) {
    session_store_->SaveSession(site, *session, mode);
  }

  unpartitioned_sessions_[SessionKey{site, session->id()}] = std::move(session);
}

void SessionServiceImpl::DeleteAllSessions(
    DeletionReason reason,
    std::optional<base::Time> created_after_time,
    std::optional<base::Time> created_before_time,
    base::RepeatingCallback<bool(const url::Origin&, const net::SchemefulSite&)>
        origin_and_site_matcher,
    base::OnceClosure completion_callback) {
  // Delete potential zombie pre-provisioned keys. Zombie keys are signing
  // keys pre-provisioned for a certain relying party that were not consumed
  // during any device bound session registration.
  // Removing zombie keys regardless of time range is fine.
  if (!origin_and_site_matcher) {
    pre_provisioned_keys_.clear();
  } else {
    std::erase_if(pre_provisioned_keys_,
                  [&](const PreProvisionedKeyEntry& key) {
                    // We only delete a pre-provisioned key if the origin and
                    // site matches the Relying Party's origin and site because
                    // the key is considered RP's data, not IdP's.
                    return origin_and_site_matcher.Run(
                        key.rp_origin, net::SchemefulSite(key.rp_origin));
                  });
  }

  if (pending_initialization_) {
    queued_operations_.push_back(base::BindOnce(
        &SessionServiceImpl::DeleteAllSessions,
        // `base::Unretained` is safe because the callback is stored in
        // `queued_operations_`, which is owned by `this`.
        base::Unretained(this), reason, created_after_time, created_before_time,
        std::move(origin_and_site_matcher), std::move(completion_callback)));
    return;
  }

  for (auto it = unpartitioned_sessions_.begin();
       it != unpartitioned_sessions_.end();) {
    auto curit = it;
    ++it;

    if (SessionMatchesFilter(curit->first.site, *curit->second,
                             created_after_time, created_before_time,
                             origin_and_site_matcher)) {
      DeleteSessionAndNotifyInternal(reason, curit, base::NullCallback());
    }
  }

  std::move(completion_callback).Run();
}

base::ScopedClosureRunner SessionServiceImpl::AddObserver(
    const GURL& url,
    base::RepeatingCallback<void(const SessionAccess&)> callback) {
  auto observer = std::make_unique<Observer>(url, callback);
  base::ScopedClosureRunner subscription(base::BindOnce(
      &SessionServiceImpl::RemoveObserver, weak_factory_.GetWeakPtr(),
      net::SchemefulSite(url), observer.get()));
  observers_by_site_[net::SchemefulSite(url)].insert(std::move(observer));
  return subscription;
}

base::CallbackListSubscription SessionServiceImpl::AddEventObserver(
    OnEventCallback callback) {
  return event_callbacks_.Add(std::move(callback));
}

void SessionServiceImpl::DeleteSessionAndNotifyInternal(
    DeletionReason reason,
    SessionServiceImpl::SessionsMap::iterator it,
    SessionService::OnAccessCallback per_request_callback) {
  LogSessionDeletionReason(reason);

  const auto& [session_key, session] = *it;

  UnblockWaitingRequests(session_key, RefreshResult::kFatalError);

  if (session_store_) {
    session_store_->DeleteSession(session_key);
  }

  NotifySessionAccess(per_request_callback,
                      SessionAccess::AccessType::kTermination, session_key,
                      *session);
  NotifyIfEventCallbackListeners([&] {
    return SessionEvent::MakeTerminationEvent(session_key.site,
                                              session_key.id.value(),
                                              /*succeeded=*/true, reason);
  });

  unpartitioned_sessions_.erase(it);
}

void SessionServiceImpl::NotifySessionAccess(
    SessionService::OnAccessCallback per_request_callback,
    SessionAccess::AccessType access_type,
    const SessionKey& session_key,
    const Session& session) {
  SessionAccess access{access_type, session_key};

  if (access_type == SessionAccess::AccessType::kTermination) {
    access.cookies.reserve(session.cookies().size());
    for (const CookieCraving& cookie : session.cookies()) {
      access.cookies.push_back(cookie.Name());
    }
  }

  if (per_request_callback) {
    per_request_callback.Run(access);
  }

  auto observers_it = observers_by_site_.find(session_key.site);
  if (observers_it == observers_by_site_.end()) {
    return;
  }

  for (const auto& observer : observers_it->second) {
    if (session.IncludesUrl(observer->url)) {
      observer->callback.Run(access);
    }
  }
}

void SessionServiceImpl::RemoveObserver(net::SchemefulSite site,
                                        Observer* observer) {
  auto observers_it = observers_by_site_.find(site);
  if (observers_it == observers_by_site_.end()) {
    return;
  }

  ObserverSet& observers = observers_it->second;

  auto it = observers.find(observer);
  if (it == observers.end()) {
    return;
  }

  observers.erase(it);

  if (observers.empty()) {
    observers_by_site_.erase(observers_it);
  }
}

SessionError::ErrorType SessionServiceImpl::OnRegistrationCompleteInternal(
    OnAccessCallback on_access_callback,
    RegistrationFetcher* fetcher,
    RegistrationResult registration_result,
    SchemefulSite site) {
  RemoveFetcher(fetcher);

  SessionError::ErrorType result =
      std::move(registration_result)
          .Visit(absl::Overload(
              [&](std::unique_ptr<Session> session) {
                CHECK(session);
                const SchemefulSite site(session->origin());
                SessionError::ErrorType success_result = SessionError::kSuccess;
                NotifyIfEventCallbackListeners([&] {
                  return SessionEvent::MakeCreationEvent(
                      site, session->id().value(), /*succeeded=*/true,
                      SessionError(success_result), session->ToDisplay());
                });
                NotifySessionAccess(on_access_callback,
                                    SessionAccess::AccessType::kCreation,
                                    SessionKey{site, session->id()}, *session);
                if (session->unexportable_key_id().has_value()) {
                  // Consume the pre-provisioned key.
                  std::erase_if(pre_provisioned_keys_,
                                [&](const PreProvisionedKeyEntry& pk) {
                                  return pk.key_id ==
                                         session->unexportable_key_id();
                                });
                }
                AddSession(site, std::move(session));
                return success_result;
              },
              [](RegistrationResult::NoSessionConfigChange)
                  -> SessionError::ErrorType {
                // This should not be returned for registrations.
                NOTREACHED();
              },
              [&](SessionError error) {
                // We failed to create a new session, so there's nothing to
                // clean up.
                SessionError::ErrorType error_type = error.type;
                NotifyIfEventCallbackListeners([&] {
                  return SessionEvent::MakeCreationEvent(
                      site, /*session_id=*/std::nullopt, /*succeeded=*/false,
                      std::move(error), /*new_session_display=*/std::nullopt);
                });
                return error_type;
              }));
  return result;
}

SessionError::ErrorType SessionServiceImpl::OnRefreshRequestCompletionInternal(
    OnAccessCallback on_access_callback,
    const SessionKey& session_key,
    RegistrationFetcher* fetcher,
    RegistrationResult registration_result) {
  RemoveFetcher(fetcher);
  CookieAndLineAccessResultList stored_cookies =
      registration_result.TakeStoredCookies();

  SessionError::ErrorType result =
      std::move(registration_result)
          .Visit(absl::Overload(
              [&](std::unique_ptr<Session> new_session) {
                // If refresh succeeded:
                // 1. update the session by adding a new session, replacing the
                //    old one
                // 2. restart the deferred requests.
                CHECK(new_session);
                CHECK_EQ(new_session->id(), session_key.id);

                Session* existing_session = GetSession(session_key);
                if (!existing_session) {
                  return SessionError::kSessionDeletedDuringRefresh;
                }

                bool is_proactive_refresh_candidate =
                    IsProactiveRefreshCandidate(*existing_session, *new_session,
                                                stored_cookies);
                std::optional<base::TimeDelta> minimum_cookie_lifetime =
                    existing_session
                        ->TakeLastProactiveRefreshOpportunityMinimumCookieLifetime();
                // Preserve the original creation date across configuration
                // refreshes. This ensures that long-standing sessions are not
                // erroneously wiped out whenever a user clears their recent
                // browsing data.
                new_session->set_creation_date(
                    existing_session->creation_date());
                // The refresh fetcher does not receive the attestation key ID
                // since AIKs are not used to sign refresh challenges. As a
                // result, the refreshed `new_session` is created without one.
                // We copy it here from `existing_session` to preserve the key
                // capability (e.g. for future on-demand key certification
                // requests) without needing to reload it from the database.
                new_session->set_unexportable_attestation_key_id(
                    existing_session->maybe_unexportable_attestation_key_id());
                SchemefulSite new_site(new_session->origin());
                // Don't bother creating the SessionDisplay if there are no
                // observers that want to know about it.
                std::optional<SessionDisplay> new_session_display =
                    event_callbacks_.empty() ? std::optional<SessionDisplay>()
                                             : new_session->ToDisplay();
                AddSession(new_site, std::move(new_session),
                           SessionStore::SaveSessionMode::kRefresh);
                // The session has been refreshed, restart the request.
                SessionError::ErrorType success_result = SessionError::kSuccess;
                UnblockWaitingRequests(session_key, RefreshResult::kRefreshed,
                                       SessionError(success_result),
                                       std::move(new_session_display),
                                       is_proactive_refresh_candidate,
                                       std::move(minimum_cookie_lifetime));
                return success_result;
              },
              [&](RegistrationResult::NoSessionConfigChange) {
                Session* existing_session = GetSession(session_key);
                if (!existing_session) {
                  return SessionError::kSessionDeletedDuringRefresh;
                }
                bool is_proactive_refresh_candidate =
                    IsProactiveRefreshCandidate(
                        *existing_session, *existing_session, stored_cookies);

                // Update the session expiry date and persist updated
                // timestamp to store.
                existing_session->RecordAccess();
                if (session_store_ &&
                    base::FeatureList::IsEnabled(
                        features::kDeviceBoundSessionsPersistExpiryOnRefresh)) {
                  session_store_->SaveSession(
                      session_key.site, *existing_session,
                      SessionStore::SaveSessionMode::kRefresh);
                }

                SessionError::ErrorType success_result = SessionError::kSuccess;
                UnblockWaitingRequests(
                    session_key, RefreshResult::kRefreshed,
                    SessionError(success_result),
                    /*new_session_display=*/std::nullopt,
                    is_proactive_refresh_candidate,
                    existing_session
                        ->TakeLastProactiveRefreshOpportunityMinimumCookieLifetime());
                return success_result;
              },
              [&](SessionError error) {
                const SessionError::ErrorType error_type = error.type;
                std::optional<DeletionReason> deletion_reason =
                    error.GetDeletionReason();
                RefreshResult refresh_result =
                    error.GetRefreshResult().value_or(
                        deletion_reason ? RefreshResult::kFatalError
                                        : RefreshResult::kUnreachable);

                UnblockWaitingRequests(session_key, refresh_result,
                                       std::move(error));

                if (deletion_reason) {
                  DeleteSessionAndNotify(*deletion_reason, session_key,
                                         on_access_callback);
                }
                return error_type;
              }));

  refresh_last_result_.insert_or_assign(session_key.site, SessionError(result));

  return result;
}

void SessionServiceImpl::RestoreSessionKey(
    const SessionKey& session_key,
    OnAccessCallback on_access_callback,
    base::OnceCallback<void(
        std::optional<unexportable_keys::UnexportableSigningKeyId>)> callback) {
  if (session_store_) {
    session_store_->RestoreSessionBindingKey(
        session_key, base::BindOnce(&SessionServiceImpl::OnSessionKeyRestored,
                                    weak_factory_.GetWeakPtr(), session_key,
                                    on_access_callback, std::move(callback)));
  } else {
    OnSessionKeyRestored(
        session_key, on_access_callback, std::move(callback),
        base::unexpected(unexportable_keys::ServiceError::kKeyNotReady));
  }
}

void SessionServiceImpl::OnSessionKeyRestored(
    const SessionKey& session_key,
    OnAccessCallback on_access_callback,
    base::OnceCallback<void(
        std::optional<unexportable_keys::UnexportableSigningKeyId>)> callback,
    Session::KeyIdOrError key_id_or_error) {
  if (!key_id_or_error.has_value()) {
    const bool is_persistent_error =
        unexportable_keys::IsPersistentError(key_id_or_error.error());
    UnblockWaitingRequests(session_key,
                           is_persistent_error
                               ? RefreshResult::kFatalError
                               : RefreshResult::kTransientSigningError);
    if (is_persistent_error) {
      DeleteSessionAndNotify(DeletionReason::kFailedToUnwrapKey, session_key,
                             on_access_callback);
    }
    std::move(callback).Run(std::nullopt);
    return;
  }

  auto* session = GetSession(session_key);
  if (!session) {
    UnblockWaitingRequests(session_key, RefreshResult::kFatalError);
    return;
  }

  session->set_unexportable_key_id(key_id_or_error);
  std::move(callback).Run(*key_id_or_error);
}

void SessionServiceImpl::StartSessionRefresh(const SessionKey& session_key,
                                             RefreshParams params) {
  const Session* session = GetSession(session_key);
  // `session` must exist because `StartSessionRefresh` is called synchronously
  // after verifying the session exists (in `OnGetCookiesForPrewarm` or
  // `DeferRequestForRefresh`).
  CHECK(session);

  const Session::KeyIdOrError& key_id = session->unexportable_key_id();
  if (!key_id.has_value()) {
    if (key_id.error() == unexportable_keys::ServiceError::kKeyNotReady) {
      SessionService::OnAccessCallback access_callback = params.access_callback;
      RestoreSessionKey(
          session_key, access_callback,
          base::BindOnce(&SessionServiceImpl::RefreshSessionInternal,
                         weak_factory_.GetWeakPtr(), std::move(params),
                         session_key));
      return;
    } else {
      UnblockWaitingRequests(session_key, RefreshResult::kFatalError);
      DeleteSessionAndNotify(DeletionReason::kFailedToRestoreKey, session_key,
                             params.access_callback);
      return;
    }
  }

  RefreshSessionInternal(std::move(params), session_key, *key_id);
}

void SessionServiceImpl::RefreshSessionInternal(
    RefreshParams params,
    const SessionKey& session_key,
    std::optional<unexportable_keys::UnexportableSigningKeyId> key_id) {
  if (!key_id) {
    UnblockWaitingRequests(session_key, RefreshResult::kFatalError);
    return;
  }

  Session* session = GetSession(session_key);
  // `session` must exist because this is either called synchronously from
  // `StartSessionRefresh` (which checks `session`), or asynchronously after
  // `OnSessionKeyRestored` which also checks `session` and returns early if
  // null.
  CHECK(session);

  if (params.trigger == RefreshTrigger::kMissingCookie) {
    auto* deferred_reqs = base::FindOrNull(deferred_requests_, session_key);
    // `deferred_reqs` must exist because `kMissingCookie` trigger implies
    // the refresh was initiated by `DeferRequestForRefresh`, which inserts
    // the session key into `deferred_requests_` before starting the refresh.
    CHECK(deferred_reqs);
    if (std::ranges::none_of(*deferred_reqs, [](const auto& req) {
          return req.triggered_refresh && !!req.request;
        })) {
      // If the original request that triggered key restoration was canceled or
      // destroyed (for example, during the asynchronous `RestoreSessionKey`
      // delay), select the next available valid request in the deferred queue
      // to act as the new triggering request for the DBSC refresh and use its
      // parameters. This prevents waiter requests from hanging indefinitely.
      auto req_it = std::ranges::find_if(
          *deferred_reqs, [](const auto& req) { return !!req.request; });
      if (req_it == deferred_reqs->end()) {
        // All deferred requests were cancelled or destroyed.
        // TODO(crbug.com/509885112): We should call UnblockWaitingRequests()
        // here to drain the queue of deferred requests because all the requests
        // have already been canceled. Use a new `RefreshResult::kCancelled` for
        // this.
        return;
      }

      req_it->triggered_refresh = true;
      params.isolation_info = req_it->request->isolation_info();
      params.site_for_cookies = req_it->request->site_for_cookies();
      params.initiator = req_it->request->initiator();
      params.access_callback =
          req_it->request->device_bound_session_access_callback();
      params.net_log = req_it->request->net_log();
    }
  }

  net::NetLogSource net_log_source_for_refresh = net::NetLogSource(
      net::NetLogSourceType::URL_REQUEST, net::NetLog::Get()->NextID());
  params.net_log.AddEventReferencingSource(
      net::NetLogEventType::DBSC_REFRESH_REQUEST, net_log_source_for_refresh);

  auto registration_param =
      RegistrationRequestParam::CreateForRefresh(*session);

  auto callback =
      base::BindOnce(&SessionServiceImpl::OnRefreshRequestCompletion,
                     weak_factory_.GetWeakPtr(), params.trigger,
                     std::move(params.access_callback), session_key);

  std::unique_ptr<RegistrationFetcher> fetcher =
      RegistrationFetcher::CreateFetcher(
          registration_param, *this, key_service_.get(), context_.get(),
          params.isolation_info, params.site_for_cookies,
          net_log_source_for_refresh, params.initiator, params.priority);
  RegistrationFetcher* fetcher_raw = fetcher.get();
  registration_fetchers_.insert(std::move(fetcher));
  fetcher_raw->StartFetchWithExistingKey(registration_param, *key_id,
                                         std::move(callback));
  // `fetcher_raw` may be deleted.
}

bool SessionServiceImpl::SigningQuotaExceeded(const SchemefulSite& site) {
  if (ignore_signing_quota_) {
    return false;
  }

  auto it = signing_times_.find(site);
  if (it == signing_times_.end()) {
    return false;
  }

  // This also discards "future" signings since `base::Time` can decrease.
  const base::Time now = base::Time::Now();
  std::erase_if(it->second, [now](base::Time time) {
    return time > now || now - time >= kSigningQuotaInterval;
  });

  size_t sign_count = it->second.size();
  if (sign_count == 0) {
    signing_times_.erase(it);
  }

  bool is_exceeded = sign_count >= kSigningQuota;
  if (auto result_it = refresh_last_result_.find(site);
      is_exceeded && result_it != refresh_last_result_.end()) {
    base::UmaHistogramEnumeration(
        "Net.DeviceBoundSessions.SigningQuotaExceededLastResult",
        result_it->second.type);
  }

  return is_exceeded;
}

void SessionServiceImpl::AddSigningOccurrence(const SchemefulSite& site) {
  signing_times_[site].push_back(base::Time::Now());
}

void SessionServiceImpl::RemoveFetcher(RegistrationFetcher* fetcher) {
  if (!fetcher) {
    return;
  }
  auto it = registration_fetchers_.find(fetcher);
  if (it == registration_fetchers_.end()) {
    return;
  }
  registration_fetchers_.erase(it);
}

void SessionServiceImpl::MaybeStartProactiveRefresh(
    SessionService::OnAccessCallback per_request_callback,
    DbscRequest& request,
    const SessionKey& session_key,
    base::TimeDelta minimum_cookie_lifetime) {
  if (minimum_cookie_lifetime > kProactiveRefreshThreshold) {
    return;
  }

  MaybeIncreaseSessionUsage(session_key, request,
                            SessionUsage::kInScopeProactiveRefreshNotPossible);

  if (deferred_requests_.find(session_key) != deferred_requests_.end()) {
    // It's not a proactive refresh if we're in the middle of a regular refresh.
    LogProactiveRefreshAttempt(
        ProactiveRefreshAttempt::kExistingDeferringRefresh);
    return;
  }

  auto* session = GetSession(session_key);
  CHECK(session);

  if (session->ShouldBackoff()) {
    LogProactiveRefreshAttempt(ProactiveRefreshAttempt::kBackoff);
    return;
  }

  if (session->attempted_proactive_refresh_since_last_success()) {
    // We only do one proactive refresh attempt before a deferral. If we
    // did not do this, every refresh due to missing cookies would be
    // skipped due to the refresh quota. Instead, we allow the refresh
    // due to missing cookies, which will communicate its reason for
    // failure in the Secure-Session-Skipped header.
    LogProactiveRefreshAttempt(
        ProactiveRefreshAttempt::kPreviousFailedProactiveRefresh);
    return;
  }

  if (!session->unexportable_key_id().has_value()) {
    // TODO(crbug.com/358137054): If we're otherwise ready for a proactive
    // refresh, we could start restoring the key. This is lower priority
    // than regular proactive refresh, since some amount of startup
    // latency is unavoidable with DBSC.
    LogProactiveRefreshAttempt(ProactiveRefreshAttempt::kMissingKey);
    return;
  }

  auto [_, inserted] = proactive_requests_.try_emplace(session_key);
  if (!inserted) {
    // Do not proactively refresh if we've already started one proactive
    // refresh.
    LogProactiveRefreshAttempt(
        ProactiveRefreshAttempt::kExistingProactiveRefresh);
    return;
  }

  MaybeIncreaseSessionUsage(session_key, request,
                            SessionUsage::kInScopeProactiveRefreshAttempted);
  NotifySessionAccess(per_request_callback, SessionAccess::AccessType::kUpdate,
                      session_key, *session);
  LogProactiveRefreshAttempt(ProactiveRefreshAttempt::kAttempted);
  StartSessionRefresh(
      session_key,
      {
          .trigger = RefreshTrigger::kProactive,
          .isolation_info = request.isolation_info(),
          .site_for_cookies = request.site_for_cookies(),
          .initiator = request.initiator(),
          .priority = unexportable_keys::BackgroundTaskPriority::kBestEffort,
          .access_callback = per_request_callback,
          .net_log = request.net_log(),
      });
}

void SessionServiceImpl::HandleResponseHeaders(
    DbscRequest& request,
    HttpResponseHeaders* headers,
    const FirstPartySetMetadata& first_party_set_metadata) {
  if (request.device_bound_session_mode() ==
      net::DeviceBoundSessionMode::kDisabled) {
    return;
  }

  const auto& request_url = request.url();

  // If response header Sec-Session-Registration is present, trigger a
  // registration request per header value to attempt to create a new session.
  std::vector<device_bound_sessions::RegistrationFetcherParam> params =
      device_bound_sessions::RegistrationFetcherParam::CreateIfValid(
          request_url, headers, restricted_sites_);
  for (auto& param : params) {
    RegisterBoundSession(request.device_bound_session_access_callback(),
                         std::move(param), request.isolation_info(),
                         request.site_for_cookies(), request.net_log(),
                         request.initiator());
  }

  // If response header Sec-Session-Challenge is present, for each header
  // value, store the challenge in advance for the next relevant refresh
  // request that gets triggered. This is to help avoid a round-trip for when
  // the next refresh request is required.
  std::vector<device_bound_sessions::SessionChallengeParam> challenge_params =
      device_bound_sessions::SessionChallengeParam::CreateIfValid(request_url,
                                                                  headers);
  for (auto& param : challenge_params) {
    SetChallengeForBoundSession(request.device_bound_session_access_callback(),
                                request, first_party_set_metadata,
                                std::move(param));
  }
}

bool CanAccessPreProvisionedKey(
    const SessionServiceImpl::CookieAccessCallback& cookie_access_cb,
    const url::Origin& provider_origin,
    const url::Origin& rp_origin) {
  return cookie_access_cb &&
         cookie_access_cb.Run({.provider_origin{provider_origin},
                               .relying_party_origin{rp_origin}});
}

bool SessionServiceImpl::CanAddPreProvisionedKey(const GURL& provider_url,
                                                 const url::Origin& rp_origin) {
  if (!CanAccessPreProvisionedKey(has_cookie_access_cb_,
                               url::Origin::Create(provider_url), rp_origin)) {
    return false;
  }

  // If we haven't reached the max keys per Identity Provider overall, we can
  // add it right away.
  if (pre_provisioned_keys_.size() <
      kMaxPreProvisionedKeysPerIdentityProvider) {
    return true;
  }

  return static_cast<size_t>(std::ranges::count(
             pre_provisioned_keys_, net::SchemefulSite(provider_url),
             &PreProvisionedKeyEntry::provider_site)) <
         kMaxPreProvisionedKeysPerIdentityProvider;
}

bool SessionServiceImpl::AddPreProvisionedKey(
    const url::Origin& rp_origin,
    std::string_view provider_key,
    const GURL& provider_url,
    unexportable_keys::UnexportableSigningKeyId key_id) {
  if (!CanAddPreProvisionedKey(provider_url, rp_origin)) {
    return false;
  }

  auto existing_key_it =
      std::ranges::find_if(pre_provisioned_keys_, [&](const auto& pk) {
        return pk.provider_url == provider_url && pk.rp_origin == rp_origin &&
               pk.provider_key == provider_key;
      });
  if (existing_key_it != pre_provisioned_keys_.end()) {
    return false;
  }

  pre_provisioned_keys_.push_back(
      {.provider_url{provider_url},
       .provider_site{net::SchemefulSite(provider_url)},
       .rp_origin{rp_origin},
       .provider_key{std::string(provider_key)},
       .key_id{key_id}});
  return true;
}

SessionErrorOr<unexportable_keys::UnexportableSigningKeyId>
SessionServiceImpl::FindPreProvisionedKey(
    const RegistrationFetcherParam& param,
    base::optional_ref<const url::Origin> original_request_initiator) {
  constexpr std::string_view kUmaMetricProviderKeyMatchOutcome =
      "Net.DeviceBoundSessions.ProviderKeyMatchOutcome";
  CHECK(param.provider_params());

  auto fail =
      [kUmaMetricProviderKeyMatchOutcome](SessionError::ErrorType error) {
        base::UmaHistogramEnumeration(kUmaMetricProviderKeyMatchOutcome, error);
        return base::unexpected(error);
      };

  if (!original_request_initiator) {
    return fail(SessionError::kInvalidPreProvisionedKeyInitiatorMissing);
  }

  const auto& provider_params = *param.provider_params();
  if (!CanAccessPreProvisionedKey(
          has_cookie_access_cb_,
          url::Origin::Create(provider_params.provider_url),
          *original_request_initiator)) {
    return fail(SessionError::kPreProvisionedKeyAccessNotGranted);
  }

  auto key_it =
      std::ranges::find_if(pre_provisioned_keys_, [&](const auto& pk) {
        return pk.provider_url == provider_params.provider_url &&
               pk.provider_key == provider_params.provider_key &&
               pk.rp_origin == *original_request_initiator;
      });
  if (key_it == pre_provisioned_keys_.end()) {
    return fail(SessionError::kPreProvisionedKeyNotFound);
  }

  base::UmaHistogramEnumeration(kUmaMetricProviderKeyMatchOutcome,
                                SessionError::kSuccess);
  return key_it->key_id;
}

}  // namespace net::device_bound_sessions
