// Copyright 2018 The Chromium Authors
// Use of this source code is governed by a BSD-style license that can be
// found in the LICENSE file.

#include "device/fido/win/type_conversions.h"

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

#include "base/compiler_specific.h"
#include "base/containers/fixed_flat_map.h"
#include "base/containers/span.h"
#include "base/logging.h"
#include "base/numerics/safe_conversions.h"
#include "base/strings/string_number_conversions.h"
#include "base/strings/utf_string_conversions.h"
#include "components/cbor/reader.h"
#include "components/device_event_log/device_event_log.h"
#include "device/fido/authenticator_get_assertion_response.h"
#include "device/fido/authenticator_make_credential_response.h"
#include "device/fido/discoverable_credential_metadata.h"
#include "device/fido/get_assertion_request_handler.h"
#include "device/fido/make_credential_request_handler.h"
#include "device/fido/opaque_attestation_statement.h"
#include "device/fido/public/fido_transport_protocol.h"
#include "device/fido/public/fido_types.h"
#include "third_party/microsoft_webauthn/src/webauthn.h"

namespace device {

namespace {

std::optional<std::vector<uint8_t>> HMACSecretOutputs(
    const WEBAUTHN_HMAC_SECRET_SALT& salt) {
  constexpr size_t kOutputLength = 32;
  if (salt.cbFirst != kOutputLength ||
      (salt.cbSecond != 0 && salt.cbSecond != kOutputLength)) {
    FIDO_LOG(ERROR) << "Incorrect HMAC output lengths: " << salt.cbFirst << " "
                    << salt.cbSecond;
    return std::nullopt;
  }

  std::vector<uint8_t> ret;
  ret.insert(ret.end(), salt.pbFirst, UNSAFE_TODO(salt.pbFirst + salt.cbFirst));
  if (salt.cbSecond == kOutputLength) {
    ret.insert(ret.end(), salt.pbSecond,
               UNSAFE_TODO(salt.pbSecond + salt.cbSecond));
  }
  return ret;
}

constexpr auto kTransportMap =
    base::MakeFixedFlatMap<DWORD, FidoTransportProtocol>(
        {{WEBAUTHN_CTAP_TRANSPORT_USB,
          FidoTransportProtocol::kUsbHumanInterfaceDevice},
         {WEBAUTHN_CTAP_TRANSPORT_NFC,
          FidoTransportProtocol::kNearFieldCommunication},
         {WEBAUTHN_CTAP_TRANSPORT_BLE,
          FidoTransportProtocol::kBluetoothLowEnergy},
         {WEBAUTHN_CTAP_TRANSPORT_INTERNAL, FidoTransportProtocol::kInternal},
         {WEBAUTHN_CTAP_TRANSPORT_HYBRID, FidoTransportProtocol::kHybrid},
         {WEBAUTHN_CTAP_TRANSPORT_SMART_CARD,
          FidoTransportProtocol::kSmartCard}});

}  // namespace

std::optional<FidoTransportProtocol> FromWinTransportsMask(
    const DWORD transport) {
  auto it = kTransportMap.find(transport);
  if (it != kTransportMap.end()) {
    return it->second;
  }

  // Ignore _TEST and possibly future others.
  return std::nullopt;
}

base::flat_set<FidoTransportProtocol> FromWinTransportsBitmask(
    const DWORD transports) {
  base::flat_set<FidoTransportProtocol> result;
  for (const auto& [mask, protocol] : kTransportMap) {
    if (transports & mask) {
      result.insert(protocol);
    }
  }
  return result;
}

uint32_t ToWinTransportsMask(
    const base::flat_set<FidoTransportProtocol>& transports) {
  uint32_t result = 0;
  for (const FidoTransportProtocol transport : transports) {
    switch (transport) {
      case FidoTransportProtocol::kUsbHumanInterfaceDevice:
        result |= WEBAUTHN_CTAP_TRANSPORT_USB;
        break;
      case FidoTransportProtocol::kNearFieldCommunication:
        result |= WEBAUTHN_CTAP_TRANSPORT_NFC;
        break;
      case FidoTransportProtocol::kBluetoothLowEnergy:
        result |= WEBAUTHN_CTAP_TRANSPORT_BLE;
        break;
      case FidoTransportProtocol::kInternal:
        result |= WEBAUTHN_CTAP_TRANSPORT_INTERNAL;
        break;
      case FidoTransportProtocol::kHybrid:
        result |= WEBAUTHN_CTAP_TRANSPORT_HYBRID;
        break;
      case FidoTransportProtocol::kSmartCard:
        result |= WEBAUTHN_CTAP_TRANSPORT_SMART_CARD;
        break;
      case FidoTransportProtocol::kDeprecatedAoa:
        // AOA is unsupported by the Windows API.
        break;
    }
  }
  return result;
}

std::optional<AuthenticatorMakeCredentialResponse>
ToAuthenticatorMakeCredentialResponse(
    const WEBAUTHN_CREDENTIAL_ATTESTATION& credential_attestation) {
  const auto authenticator_data_span =
      ToAuthenticatorDataSpan(credential_attestation);
  auto authenticator_data =
      AuthenticatorData::DecodeAuthenticatorData(authenticator_data_span);
  if (!authenticator_data) {
    DLOG(ERROR) << "DecodeAuthenticatorData failed: "
                << base::HexEncode(authenticator_data_span);
    return std::nullopt;
  }
  const auto attestation_span = ToAttestationSpan(credential_attestation);
  std::optional<cbor::Value> cbor_attestation_statement =
      cbor::Reader::Read(attestation_span);
  if (!cbor_attestation_statement || !cbor_attestation_statement->is_map()) {
    DLOG(ERROR) << "CBOR decoding attestation statement failed: "
                << base::HexEncode(attestation_span);
    return std::nullopt;
  }

  std::optional<FidoTransportProtocol> transport_used;
  if (credential_attestation.dwVersion >=
      WEBAUTHN_CREDENTIAL_ATTESTATION_VERSION_3) {
    // dwUsedTransport should have exactly one of the
    // WEBAUTHN_CTAP_TRANSPORT_* values set.
    transport_used =
        FromWinTransportsMask(credential_attestation.dwUsedTransport);
  }

  AuthenticatorMakeCredentialResponse ret(
      transport_used,
      AttestationObject(
          std::move(*authenticator_data),
          std::make_unique<OpaqueAttestationStatement>(
              base::WideToUTF8(credential_attestation.pwszFormatType),
              std::move(*cbor_attestation_statement))));
  if (credential_attestation.dwVersion >=
      WEBAUTHN_CREDENTIAL_ATTESTATION_VERSION_8) {
    ret.transports =
        FromWinTransportsBitmask(credential_attestation.dwTransports);
  } else if (transport_used == FidoTransportProtocol::kInternal) {
    // Before webauthn.dll version 9, Windows would only enumerate platform
    // credentials. These credentials can't be used from other devices, so we
    // can fill in the authenticator supported transports.
    ret.transports = {*transport_used};
  }

  if (credential_attestation.dwVersion >=
      WEBAUTHN_CREDENTIAL_ATTESTATION_VERSION_4) {
    ret.enterprise_attestation_returned = credential_attestation.bEpAtt;
    ret.is_resident_key = credential_attestation.bResidentKey;
    if (credential_attestation.bLargeBlobSupported) {
      ret.large_blob_type = LargeBlobSupportType::kBespoke;
    }
  }

  if (credential_attestation.dwVersion >=
      WEBAUTHN_CREDENTIAL_ATTESTATION_VERSION_5) {
    ret.prf_enabled = credential_attestation.bPrfEnabled;
  }

  if (credential_attestation.dwVersion >=
          WEBAUTHN_CREDENTIAL_ATTESTATION_VERSION_7 &&
      credential_attestation.pHmacSecret) {
    ret.prf_results = HMACSecretOutputs(*credential_attestation.pHmacSecret);
  }

  return ret;
}

std::optional<AuthenticatorGetAssertionResponse>
ToAuthenticatorGetAssertionResponse(
    const WEBAUTHN_ASSERTION& assertion,
    const CtapGetAssertionOptions& request_options) {
  const auto authenticator_data_span = ToAuthenticatorDataSpan(assertion);
  auto authenticator_data =
      AuthenticatorData::DecodeAuthenticatorData(authenticator_data_span);
  if (!authenticator_data) {
    DLOG(ERROR) << "DecodeAuthenticatorData failed: "
                << base::HexEncode(authenticator_data_span);
    return std::nullopt;
  }
  std::optional<FidoTransportProtocol> transport_used =
      assertion.dwVersion >= WEBAUTHN_ASSERTION_VERSION_4
          ? FromWinTransportsMask(assertion.dwUsedTransport)
          : std::nullopt;
  AuthenticatorGetAssertionResponse response(
      std::move(*authenticator_data),
      std::vector<uint8_t>(
          assertion.pbSignature,
          UNSAFE_TODO(assertion.pbSignature + assertion.cbSignature)),
      transport_used);
  response.credential = PublicKeyCredentialDescriptor(
      CredentialType::kPublicKey,
      std::vector<uint8_t>(
          assertion.Credential.pbId,
          UNSAFE_TODO(assertion.Credential.pbId + assertion.Credential.cbId)));
  if (assertion.cbUserId > 0) {
    response.user_entity = PublicKeyCredentialUserEntity(std::vector<uint8_t>(
        assertion.pbUserId,
        UNSAFE_TODO(assertion.pbUserId + assertion.cbUserId)));
  }
  if (assertion.dwVersion >= WEBAUTHN_ASSERTION_VERSION_2 &&
      assertion.dwCredLargeBlobStatus ==
          WEBAUTHN_CRED_LARGE_BLOB_STATUS_SUCCESS) {
    if (request_options.large_blob_read) {
      response.large_blob = std::vector<uint8_t>(
          assertion.pbCredLargeBlob,
          UNSAFE_TODO(assertion.pbCredLargeBlob + assertion.cbCredLargeBlob));
    } else if (request_options.large_blob_write) {
      response.large_blob_written = true;
    }
  }
  if (assertion.dwVersion >= WEBAUTHN_ASSERTION_VERSION_3 &&
      assertion.pHmacSecret) {
    response.hmac_secret = HMACSecretOutputs(*assertion.pHmacSecret);
  }
  return response;
}

uint32_t ToWinUserVerificationRequirement(
    UserVerificationRequirement user_verification_requirement) {
  switch (user_verification_requirement) {
    case UserVerificationRequirement::kRequired:
      return WEBAUTHN_USER_VERIFICATION_REQUIREMENT_REQUIRED;
    case UserVerificationRequirement::kPreferred:
      return WEBAUTHN_USER_VERIFICATION_REQUIREMENT_PREFERRED;
    case UserVerificationRequirement::kDiscouraged:
      return WEBAUTHN_USER_VERIFICATION_REQUIREMENT_DISCOURAGED;
  }
  NOTREACHED();
}

uint32_t ToWinAuthenticatorAttachment(
    AuthenticatorAttachment authenticator_attachment) {
  switch (authenticator_attachment) {
    case AuthenticatorAttachment::kAny:
      return WEBAUTHN_AUTHENTICATOR_ATTACHMENT_ANY;
    case AuthenticatorAttachment::kPlatform:
      return WEBAUTHN_AUTHENTICATOR_ATTACHMENT_PLATFORM;
    case AuthenticatorAttachment::kCrossPlatform:
      return WEBAUTHN_AUTHENTICATOR_ATTACHMENT_CROSS_PLATFORM;
  }
  NOTREACHED();
}

std::vector<WEBAUTHN_CREDENTIAL> ToWinCredentialVector(
    const std::vector<PublicKeyCredentialDescriptor>* credentials) {
  std::vector<WEBAUTHN_CREDENTIAL> result;
  for (const auto& credential : *credentials) {
    if (credential.credential_type != CredentialType::kPublicKey) {
      continue;
    }
    result.push_back(WEBAUTHN_CREDENTIAL{
        WEBAUTHN_CREDENTIAL_CURRENT_VERSION,
        base::checked_cast<DWORD>(credential.id.size()),
        const_cast<unsigned char*>(credential.id.data()),
        WEBAUTHN_CREDENTIAL_TYPE_PUBLIC_KEY,
    });
  }
  return result;
}

std::vector<WEBAUTHN_CREDENTIAL_EX> ToWinCredentialExVector(
    const std::vector<PublicKeyCredentialDescriptor>* credentials) {
  std::vector<WEBAUTHN_CREDENTIAL_EX> result;
  for (const auto& credential : *credentials) {
    if (credential.credential_type != CredentialType::kPublicKey) {
      continue;
    }
    result.push_back(
        WEBAUTHN_CREDENTIAL_EX{WEBAUTHN_CREDENTIAL_EX_CURRENT_VERSION,
                               base::checked_cast<DWORD>(credential.id.size()),
                               const_cast<unsigned char*>(credential.id.data()),
                               WEBAUTHN_CREDENTIAL_TYPE_PUBLIC_KEY,
                               ToWinTransportsMask(credential.transports)});
  }
  return result;
}

uint32_t ToWinLargeBlobSupport(LargeBlobSupport large_blob_support) {
  switch (large_blob_support) {
    case LargeBlobSupport::kNotRequested:
      return WEBAUTHN_LARGE_BLOB_SUPPORT_NONE;
    case LargeBlobSupport::kPreferred:
      return WEBAUTHN_LARGE_BLOB_SUPPORT_PREFERRED;
    case LargeBlobSupport::kRequired:
      return WEBAUTHN_LARGE_BLOB_SUPPORT_REQUIRED;
  }
}

COMPONENT_EXPORT(DEVICE_FIDO)
MakeCredentialStatus WinErrorNameToMakeCredentialStatus(
    std::u16string_view error_name) {
  // See WebAuthNGetErrorName in <webauthn.h> for these string literals.
  constexpr auto kResponseCodeMap =
      base::MakeFixedFlatMap<std::u16string_view, MakeCredentialStatus>({
          {u"Success", MakeCredentialStatus::kSuccess},
          {u"InvalidStateError",
           MakeCredentialStatus::kUserConsentButCredentialExcluded},
          {u"ConstraintError",
           MakeCredentialStatus::kAuthenticatorResponseInvalid},
          {u"NotSupportedError",
           MakeCredentialStatus::kAuthenticatorResponseInvalid},
          {u"NotAllowedError", MakeCredentialStatus::kWinNotAllowedError},
          {u"UnknownError",
           MakeCredentialStatus::kAuthenticatorResponseInvalid},
      });
  const auto it = kResponseCodeMap.find(error_name);
  if (it == kResponseCodeMap.end()) {
    FIDO_LOG(ERROR) << "Unexpected error name: " << error_name;
    return MakeCredentialStatus::kAuthenticatorResponseInvalid;
  }
  return it->second;
}

GetAssertionStatus WinErrorNameToGetAssertionStatus(
    std::u16string_view error_name) {
  // See WebAuthNGetErrorName in <webauthn.h> for these string literals.
  //
  // "NotAllowedError" indicates the user cancelled, there was no matching
  // credential, or a timeout. Other errors indicate that either the
  // request was rejected or there was an error processing it.
  constexpr auto kResponseCodeMap = base::MakeFixedFlatMap<std::u16string_view,
                                                           GetAssertionStatus>({
      {u"Success", GetAssertionStatus::kSuccess},
      {u"InvalidStateError", GetAssertionStatus::kAuthenticatorResponseInvalid},
      {u"ConstraintError", GetAssertionStatus ::kAuthenticatorResponseInvalid},
      {u"NotSupportedError", GetAssertionStatus::kAuthenticatorResponseInvalid},
      {u"NotAllowedError", GetAssertionStatus::kWinNotAllowedError},
      {u"UnknownError", GetAssertionStatus::kAuthenticatorResponseInvalid},
  });
  const auto it = kResponseCodeMap.find(error_name);
  if (it == kResponseCodeMap.end()) {
    FIDO_LOG(ERROR) << "Unexpected error name: " << error_name;
    return GetAssertionStatus::kAuthenticatorResponseInvalid;
  }
  return it->second;
}

uint32_t ToWinAttestationConveyancePreference(
    const AttestationConveyancePreference& value,
    int api_version) {
  switch (value) {
    case AttestationConveyancePreference::kNone:
      return WEBAUTHN_ATTESTATION_CONVEYANCE_PREFERENCE_NONE;
    case AttestationConveyancePreference::kIndirect:
      return WEBAUTHN_ATTESTATION_CONVEYANCE_PREFERENCE_DIRECT;
    case AttestationConveyancePreference::kDirect:
      return WEBAUTHN_ATTESTATION_CONVEYANCE_PREFERENCE_DIRECT;
    case AttestationConveyancePreference::kEnterpriseIfRPListedOnAuthenticator:
    case AttestationConveyancePreference::kEnterpriseApprovedByBrowser:
      // Enterprise attestation is supported in API version 3.
      return api_version >= 3
                 ? WEBAUTHN_ATTESTATION_CONVEYANCE_PREFERENCE_DIRECT
                 : WEBAUTHN_ATTESTATION_CONVEYANCE_PREFERENCE_NONE;
  }
  NOTREACHED();
}

std::vector<DiscoverableCredentialMetadata>
WinCredentialDetailsListToCredentialMetadata(
    const WEBAUTHN_CREDENTIAL_DETAILS_LIST& credentials) {
  std::vector<DiscoverableCredentialMetadata> result;
  for (size_t i = 0; i < credentials.cCredentialDetails; ++i) {
    WEBAUTHN_CREDENTIAL_DETAILS* credential =
        UNSAFE_TODO(credentials.ppCredentialDetails[i]);
    WEBAUTHN_USER_ENTITY_INFORMATION* user = credential->pUserInformation;
    WEBAUTHN_RP_ENTITY_INFORMATION* rp = credential->pRpInformation;
    DiscoverableCredentialMetadata metadata(
        AuthenticatorType::kWinNative, base::WideToUTF8(rp->pwszId),
        std::vector<uint8_t>(credential->pbCredentialID,
                             UNSAFE_TODO(credential->pbCredentialID +
                                         credential->cbCredentialID)),
        PublicKeyCredentialUserEntity(
            std::vector<uint8_t>(user->pbId,
                                 UNSAFE_TODO(user->pbId + user->cbId)),
            user->pwszName
                ? std::make_optional(base::WideToUTF8(user->pwszName))
                : std::nullopt,
            user->pwszDisplayName
                ? std::make_optional(base::WideToUTF8(user->pwszDisplayName))
                : std::nullopt),
        credential->dwVersion >= WEBAUTHN_CREDENTIAL_DETAILS_VERSION_3 &&
                credential->pwszAuthenticatorName
            ? std::make_optional(
                  base::WideToUTF8(credential->pwszAuthenticatorName))
            : std::nullopt);
    metadata.system_created = !credential->bRemovable;
    if (credential->dwVersion >= WEBAUTHN_CREDENTIAL_DETAILS_VERSION_4) {
      metadata.transports = FromWinTransportsBitmask(credential->dwTransports);
    }
    result.push_back(std::move(metadata));
  }
  return result;
}

std::vector<const wchar_t*> ToWinCredentialHints(
    base::span<const blink::mojom::Hint> hints) {
  std::vector<const wchar_t*> ret;
  ret.reserve(hints.size());
  for (const blink::mojom::Hint& hint : hints) {
    switch (hint) {
      case blink::mojom::Hint::SECURITY_KEY:
        ret.emplace_back(WEBAUTHN_CREDENTIAL_HINT_SECURITY_KEY);
        break;
      case blink::mojom::Hint::CLIENT_DEVICE:
        ret.emplace_back(WEBAUTHN_CREDENTIAL_HINT_CLIENT_DEVICE);
        break;
      case blink::mojom::Hint::HYBRID:
        ret.emplace_back(WEBAUTHN_CREDENTIAL_HINT_HYBRID);
        break;
    }
  }
  return ret;
}

base::span<const uint8_t> ToAuthenticatorDataSpan(
    const WEBAUTHN_ASSERTION& in) {
  // SAFETY: The size of `in.pbAuthenticatorData` must be
  // `in.cbAuthenticatorData`.
  return UNSAFE_BUFFERS(base::span<const uint8_t>(in.pbAuthenticatorData,
                                                  in.cbAuthenticatorData));
}

base::span<const uint8_t> ToAuthenticatorDataSpan(
    const WEBAUTHN_CREDENTIAL_ATTESTATION& in) {
  // SAFETY: The size of `in.pbAuthenticatorData` must be
  // `in.cbAuthenticatorData`.
  return UNSAFE_BUFFERS(base::span<const uint8_t>(in.pbAuthenticatorData,
                                                  in.cbAuthenticatorData));
}

base::span<const uint8_t> ToUserIdSpan(const WEBAUTHN_ASSERTION& in) {
  // SAFETY: The size of `in.pbUserId` must be `in.cbUserId`.
  return UNSAFE_BUFFERS(base::span<const uint8_t>(in.pbUserId, in.cbUserId));
}

base::span<const uint8_t> ToCredentialIdSpan(
    const WEBAUTHN_CREDENTIAL_ATTESTATION& in) {
  // SAFETY: The size of `in.pbCredentialId` must be `in.cbCredentialId`.
  return UNSAFE_BUFFERS(
      base::span<const uint8_t>(in.pbCredentialId, in.cbCredentialId));
}

base::span<const uint8_t> ToAttestationSpan(
    const WEBAUTHN_CREDENTIAL_ATTESTATION& in) {
  // SAFETY: The size of `in.pbAttestation` must be `in.cbAttestation`.
  return UNSAFE_BUFFERS(
      base::span<const uint8_t>(in.pbAttestation, in.cbAttestation));
}

base::span<const uint8_t> ToAttestationObjectSpan(
    const WEBAUTHN_CREDENTIAL_ATTESTATION& in) {
  // SAFETY: The size of `in.pbAttestationObject` must be
  // `in.cbAttestationObject`.
  return UNSAFE_BUFFERS(base::span<const uint8_t>(in.pbAttestationObject,
                                                  in.cbAttestationObject));
}

base::span<const uint8_t> ToIdSpan(const WEBAUTHN_CREDENTIAL& in) {
  // SAFETY: The size of `in.pbId` must be `in.cbId`.
  return UNSAFE_BUFFERS(base::span<const uint8_t>(in.pbId, in.cbId));
}

base::span<const uint8_t> ToIdSpan(const WEBAUTHN_CREDENTIAL_EX& in) {
  // SAFETY: The size of `in.pbId` must be `in.cbId`.
  return UNSAFE_BUFFERS(base::span<const uint8_t>(in.pbId, in.cbId));
}

base::span<const uint8_t> ToIdSpan(const WEBAUTHN_USER_ENTITY_INFORMATION& in) {
  // SAFETY: The size of `in.pbId` must be `in.cbId`.
  return UNSAFE_BUFFERS(base::span<const uint8_t>(in.pbId, in.cbId));
}

base::span<const uint8_t> ToExtensionSpan(const WEBAUTHN_EXTENSION& in) {
  // SAFETY: The size of `in.pvExtension` must be `in.cbExtension`.
  return UNSAFE_BUFFERS(base::span<const uint8_t>(
      reinterpret_cast<const uint8_t*>(in.pvExtension), in.cbExtension));
}

base::span<const uint8_t> ToCredIdSpan(
    const WEBAUTHN_CRED_WITH_HMAC_SECRET_SALT& in) {
  // SAFETY: The size of `in.pbCredID` must be `in.cbCredID`.
  return UNSAFE_BUFFERS(base::span<const uint8_t>(in.pbCredID, in.cbCredID));
}

}  // namespace device
