// 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 <memory>
#include <utility>

#include "testing/gtest/include/gtest/gtest.h"
#include "third_party/blink/public/mojom/frame/user_activation_notification_type.mojom-blink.h"
#include "third_party/blink/public/mojom/webid/digital_identity_request.mojom.h"
#include "third_party/blink/renderer/bindings/core/v8/script_promise_resolver.h"
#include "third_party/blink/renderer/bindings/core/v8/script_promise_tester.h"
#include "third_party/blink/renderer/bindings/core/v8/script_value.h"
#include "third_party/blink/renderer/bindings/core/v8/v8_binding_for_core.h"
#include "third_party/blink/renderer/bindings/core/v8/v8_binding_for_testing.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_credential_creation_options.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_credential_request_options.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_digital_credential_create_request.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_digital_credential_creation_options.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_digital_credential_get_request.h"
#include "third_party/blink/renderer/bindings/modules/v8/v8_digital_credential_request_options.h"
#include "third_party/blink/renderer/core/dom/document.h"
#include "third_party/blink/renderer/core/frame/local_dom_window.h"
#include "third_party/blink/renderer/core/frame/local_frame.h"
#include "third_party/blink/renderer/core/testing/page_test_base.h"
#include "third_party/blink/renderer/modules/credentialmanagement/credential.h"
#include "third_party/blink/renderer/modules/credentialmanagement/digital_credential.h"
#include "third_party/blink/renderer/modules/credentialmanagement/digital_identity_credential.h"
#include "third_party/blink/renderer/platform/heap/collection_support/heap_vector.h"
#include "third_party/blink/renderer/platform/heap/garbage_collected.h"
#include "third_party/blink/renderer/platform/testing/runtime_enabled_features_test_helpers.h"
#include "third_party/blink/renderer/platform/testing/unit_test_helpers.h"
#include "third_party/blink/renderer/platform/weborigin/kurl.h"

namespace blink {

namespace {

// Mock mojom::DigitalIdentityRequest which succeeds and returns "token".
class MockDigitalIdentityRequest : public mojom::DigitalIdentityRequest {
 public:
  MockDigitalIdentityRequest() = default;

  MockDigitalIdentityRequest(const MockDigitalIdentityRequest&) = delete;
  MockDigitalIdentityRequest& operator=(const MockDigitalIdentityRequest&) =
      delete;

  void Bind(mojo::PendingReceiver<mojom::DigitalIdentityRequest> receiver) {
    receiver_.Bind(std::move(receiver));
  }
  void Get(std::vector<blink::mojom::DigitalCredentialGetRequestPtr> requests,
           GetCallback callback) override {
    std::move(callback).Run(mojom::RequestDigitalIdentityStatus::kSuccess,
                            "protocol", base::Value("token"));
  }

  void Create(
      std::vector<blink::mojom::DigitalCredentialCreateRequestPtr> requests,
      CreateCallback callback) override {
    std::move(callback).Run(mojom::RequestDigitalIdentityStatus::kSuccess,
                            "protocol", base::Value("token"));
  }

  void Abort() override {}

 private:
  mojo::Receiver<mojom::DigitalIdentityRequest> receiver_{this};
};

CredentialRequestOptions* CreateGetOptionsWithRequests(
    const HeapVector<Member<DigitalCredentialGetRequest>>& requests) {
  DigitalCredentialRequestOptions* digital_credential_request =
      DigitalCredentialRequestOptions::Create();
  digital_credential_request->setRequests(requests);
  CredentialRequestOptions* options = CredentialRequestOptions::Create();
  options->setDigital(digital_credential_request);
  return options;
}

CredentialRequestOptions* CreateOptionsWithProtocol(ScriptState* script_state,
                                                    const String& protocol) {
  DigitalCredentialGetRequest* request = DigitalCredentialGetRequest::Create();
  request->setProtocol(protocol);
  v8::Local<v8::Object> request_data =
      v8::Object::New(script_state->GetIsolate());

  request->setData(ScriptObject(script_state->GetIsolate(), request_data));
  HeapVector<Member<DigitalCredentialGetRequest>> requests;
  requests.push_back(request);
  return CreateGetOptionsWithRequests(requests);
}

CredentialCreationOptions* CreateCreateOptionsWithRequests(
    const HeapVector<Member<DigitalCredentialCreateRequest>>& requests) {
  DigitalCredentialCreationOptions* digital_credential_request =
      DigitalCredentialCreationOptions::Create();
  digital_credential_request->setRequests(requests);
  CredentialCreationOptions* options = CredentialCreationOptions::Create();
  options->setDigital(digital_credential_request);
  return options;
}

CredentialCreationOptions* CreateCreateOptionsWithProtocol(
    ScriptState* script_state,
    const String& protocol) {
  DigitalCredentialCreateRequest* request =
      DigitalCredentialCreateRequest::Create();
  request->setProtocol(protocol);
  v8::Local<v8::Object> request_data =
      v8::Object::New(script_state->GetIsolate());

  request->setData(ScriptObject(script_state->GetIsolate(), request_data));
  HeapVector<Member<DigitalCredentialCreateRequest>> requests;
  requests.push_back(request);
  return CreateCreateOptionsWithRequests(requests);
}

}  // namespace

class DigitalIdentityCredentialProtocolTest : public PageTestBase {
 public:
  DigitalIdentityCredentialProtocolTest() = default;
  ~DigitalIdentityCredentialProtocolTest() override = default;

  void SetUp() override {
    EnablePlatform();
    PageTestBase::SetUp();

    NavigateTo(KURL("https://example.test"));

    mock_request_ = std::make_unique<MockDigitalIdentityRequest>();
    GetFrame().DomWindow()->GetBrowserInterfaceBroker().SetBinderForTesting(
        mojom::DigitalIdentityRequest::Name_,
        BindRepeating(
            [](MockDigitalIdentityRequest* mock_request_ptr,
               mojo::ScopedMessagePipeHandle handle) {
              mock_request_ptr->Bind(
                  mojo::PendingReceiver<mojom::DigitalIdentityRequest>(
                      std::move(handle)));
            },
            Unretained(mock_request_.get())));
  }

  void TearDown() override {
    GetFrame().DomWindow()->GetBrowserInterfaceBroker().SetBinderForTesting(
        mojom::DigitalIdentityRequest::Name_, {});
    PageTestBase::TearDown();
  }

 protected:
  std::unique_ptr<MockDigitalIdentityRequest> mock_request_;
};

TEST_F(DigitalIdentityCredentialProtocolTest, DiscoverProtocolUseCounters) {
  ScopedWebIdentityDigitalCredentialsForTest scoped_digital_credentials(
      /*enabled=*/true);

  struct TestCase {
    String protocol;
    mojom::WebFeature feature;
  };

  TestCase test_cases[] = {
      {"openid4vp-v1-unsigned",
       mojom::WebFeature::kDigitalCredentialsProtocolOpenId4VpUnsigned},
      {"openid4vp-v1-signed",
       mojom::WebFeature::kDigitalCredentialsProtocolOpenId4VpSigned},
      {"openid4vp-v1-multisigned",
       mojom::WebFeature::kDigitalCredentialsProtocolOpenId4VpMultisigned},
      {"org-iso-mdoc",
       mojom::WebFeature::kDigitalCredentialsProtocolOrgIsoMdoc},
  };

  for (const auto& test_case : test_cases) {
    ScriptState* script_state = ToScriptStateForMainWorld(&GetFrame());
    ScriptState::Scope scope(script_state);
    auto* resolver =
        MakeGarbageCollected<ScriptPromiseResolver<IDLNullable<Credential>>>(
            script_state);

    GetFrame().NotifyUserActivation(
        mojom::blink::UserActivationNotificationType::kTest);
    DiscoverDigitalIdentityCredentialFromExternalSource(
        resolver, *CreateOptionsWithProtocol(script_state, test_case.protocol));

    test::RunPendingTasks();

    EXPECT_TRUE(GetDocument().IsUseCounted(test_case.feature))
        << "Feature not counted for protocol: " << test_case.protocol;
  }
}

TEST_F(DigitalIdentityCredentialProtocolTest, CreateProtocolUseCounters) {
  ScopedWebIdentityDigitalCredentialsCreationForTest
      scoped_digital_credentials_creation(
          /*enabled=*/true);

  struct TestCase {
    String protocol;
    mojom::WebFeature feature;
  };

  TestCase test_cases[] = {
      {"openid4vci", mojom::WebFeature::kDigitalCredentialsProtocolOpenId4Vci},
      {"openid4vci-v1",
       mojom::WebFeature::kDigitalCredentialsProtocolOpenId4VciV1},
  };

  for (const auto& test_case : test_cases) {
    ScriptState* script_state = ToScriptStateForMainWorld(&GetFrame());
    ScriptState::Scope scope(script_state);
    auto* resolver =
        MakeGarbageCollected<ScriptPromiseResolver<IDLNullable<Credential>>>(
            script_state);

    GetFrame().NotifyUserActivation(
        mojom::blink::UserActivationNotificationType::kTest);
    CreateDigitalIdentityCredentialInExternalSource(
        resolver,
        *CreateCreateOptionsWithProtocol(script_state, test_case.protocol));

    test::RunPendingTasks();

    EXPECT_TRUE(GetDocument().IsUseCounted(test_case.feature))
        << "Feature not counted for protocol: " << test_case.protocol;
  }
}

TEST_F(DigitalIdentityCredentialProtocolTest,
       DiscoverProtocolUseCountersMultipleRequests) {
  ScopedWebIdentityDigitalCredentialsForTest scoped_digital_credentials(
      /*enabled=*/true);

  ScriptState* script_state = ToScriptStateForMainWorld(&GetFrame());
  ScriptState::Scope scope(script_state);
  auto* resolver =
      MakeGarbageCollected<ScriptPromiseResolver<IDLNullable<Credential>>>(
          script_state);

  HeapVector<Member<DigitalCredentialGetRequest>> requests;
  {
    DigitalCredentialGetRequest* request =
        DigitalCredentialGetRequest::Create();
    request->setProtocol("org-iso-mdoc");
    v8::Local<v8::Object> request_data =
        v8::Object::New(script_state->GetIsolate());
    request->setData(ScriptObject(script_state->GetIsolate(), request_data));
    requests.push_back(request);
  }
  {
    DigitalCredentialGetRequest* request =
        DigitalCredentialGetRequest::Create();
    request->setProtocol("openid4vp-v1-unsigned");
    v8::Local<v8::Object> request_data =
        v8::Object::New(script_state->GetIsolate());
    request->setData(ScriptObject(script_state->GetIsolate(), request_data));
    requests.push_back(request);
  }

  GetFrame().NotifyUserActivation(
      mojom::blink::UserActivationNotificationType::kTest);
  DiscoverDigitalIdentityCredentialFromExternalSource(
      resolver, *CreateGetOptionsWithRequests(requests));

  test::RunPendingTasks();

  EXPECT_TRUE(GetDocument().IsUseCounted(
      mojom::WebFeature::kDigitalCredentialsProtocolOrgIsoMdoc));
  EXPECT_TRUE(GetDocument().IsUseCounted(
      mojom::WebFeature::kDigitalCredentialsProtocolOpenId4VpUnsigned));
}

TEST_F(DigitalIdentityCredentialProtocolTest,
       DiscoverProtocolUseCountersUnknownProtocol) {
  ScopedWebIdentityDigitalCredentialsForTest scoped_digital_credentials(
      /*enabled=*/true);

  ScriptState* script_state = ToScriptStateForMainWorld(&GetFrame());
  ScriptState::Scope scope(script_state);
  auto* resolver =
      MakeGarbageCollected<ScriptPromiseResolver<IDLNullable<Credential>>>(
          script_state);

  GetFrame().NotifyUserActivation(
      mojom::blink::UserActivationNotificationType::kTest);
  DiscoverDigitalIdentityCredentialFromExternalSource(
      resolver, *CreateOptionsWithProtocol(script_state, "unknown-protocol"));

  test::RunPendingTasks();

  EXPECT_TRUE(GetDocument().IsUseCounted(
      mojom::WebFeature::kDigitalCredentialsProtocolUnknown));
}

TEST_F(DigitalIdentityCredentialProtocolTest,
       CreateProtocolUseCountersUnknownProtocol) {
  ScopedWebIdentityDigitalCredentialsCreationForTest
      scoped_digital_credentials_creation(
          /*enabled=*/true);

  ScriptState* script_state = ToScriptStateForMainWorld(&GetFrame());
  ScriptState::Scope scope(script_state);
  auto* resolver =
      MakeGarbageCollected<ScriptPromiseResolver<IDLNullable<Credential>>>(
          script_state);

  GetFrame().NotifyUserActivation(
      mojom::blink::UserActivationNotificationType::kTest);
  CreateDigitalIdentityCredentialInExternalSource(
      resolver,
      *CreateCreateOptionsWithProtocol(script_state, "unknown-protocol"));

  test::RunPendingTasks();

  EXPECT_TRUE(GetDocument().IsUseCounted(
      mojom::WebFeature::kDigitalCredentialsProtocolUnknown));
}

class DigitalIdentityCredentialProtocolFilterTest
    : public DigitalIdentityCredentialProtocolTest {
 protected:
  void ExpectDiscoverRejects(const String& protocol) {
    ScriptState* script_state = ToScriptStateForMainWorld(&GetFrame());
    ScriptState::Scope scope(script_state);
    auto* resolver =
        MakeGarbageCollected<ScriptPromiseResolver<IDLNullable<Credential>>>(
            script_state);

    GetFrame().NotifyUserActivation(
        mojom::blink::UserActivationNotificationType::kTest);
    DiscoverDigitalIdentityCredentialFromExternalSource(
        resolver, *CreateOptionsWithProtocol(script_state, protocol));

    ScriptPromiseTester tester(script_state, resolver->Promise());
    tester.WaitUntilSettled();

    EXPECT_TRUE(tester.IsRejected())
        << "Expected Discover to reject protocol: " << protocol;
    EXPECT_TRUE(GetDocument().IsUseCounted(
        mojom::WebFeature::kDigitalCredentialsProtocolUnknown));
  }

  void ExpectCreateRejects(const String& protocol) {
    ScriptState* script_state = ToScriptStateForMainWorld(&GetFrame());
    ScriptState::Scope scope(script_state);
    auto* resolver =
        MakeGarbageCollected<ScriptPromiseResolver<IDLNullable<Credential>>>(
            script_state);

    GetFrame().NotifyUserActivation(
        mojom::blink::UserActivationNotificationType::kTest);
    CreateDigitalIdentityCredentialInExternalSource(
        resolver, *CreateCreateOptionsWithProtocol(script_state, protocol));

    ScriptPromiseTester tester(script_state, resolver->Promise());
    tester.WaitUntilSettled();

    EXPECT_TRUE(tester.IsRejected())
        << "Expected Create to reject protocol: " << protocol;
    EXPECT_TRUE(GetDocument().IsUseCounted(
        mojom::WebFeature::kDigitalCredentialsProtocolUnknown));
  }
};

TEST_F(DigitalIdentityCredentialProtocolFilterTest,
       DiscoverRejectsInvalidProtocolsWhenFeatureEnabled) {
  ScopedWebIdentityDigitalCredentialsForTest scoped_digital_credentials(
      /*enabled=*/true);
  ScopedDigitalCredentialsProtocolFilterForTest scoped_feature(
      /*enabled=*/true);

  // "unknown-protocol" is completely unknown.
  // "openid4vci" is an issuance protocol, which is invalid for Get/Discover.
  const String kInvalidProtocols[] = {"unknown-protocol", "openid4vci"};

  for (const String& protocol : kInvalidProtocols) {
    ExpectDiscoverRejects(protocol);
  }
}

TEST_F(DigitalIdentityCredentialProtocolFilterTest,
       CreateRejectsInvalidProtocolsWhenFeatureEnabled) {
  ScopedWebIdentityDigitalCredentialsCreationForTest
      scoped_digital_credentials_creation(/*enabled=*/true);
  ScopedDigitalCredentialsProtocolFilterForTest scoped_feature(
      /*enabled=*/true);

  // "unknown-protocol" is completely unknown.
  // "org-iso-mdoc" is a presentation protocol, which is invalid for Create.
  const String kInvalidProtocols[] = {"unknown-protocol", "org-iso-mdoc"};

  for (const String& protocol : kInvalidProtocols) {
    ExpectCreateRejects(protocol);
  }
}

TEST_F(DigitalIdentityCredentialProtocolFilterTest,
       UserAgentAllowsProtocolFilter) {
  ScriptState* script_state = ToScriptStateForMainWorld(&GetFrame());
  ScriptState::Scope scope(script_state);

  {
    ScopedDigitalCredentialsProtocolFilterForTest scoped_feature(
        /*enabled=*/false);

    // When feature is disabled, any format-compliant protocol is allowed.
    const String kAllowed[] = {"openid4vp", "openid4vp-v1-unsigned",
                               "some-random-protocol"};
    for (const String& protocol : kAllowed) {
      EXPECT_TRUE(
          DigitalCredential::userAgentAllowsProtocol(script_state, protocol));
    }

    EXPECT_FALSE(DigitalCredential::userAgentAllowsProtocol(
        script_state, "INVALID_PROTOCOL"));
  }

  {
    ScopedDigitalCredentialsProtocolFilterForTest scoped_feature(
        /*enabled=*/true);

    // When feature is enabled, only known supported protocols are allowed.
    const String kAllowed[] = {"openid4vp-v1-unsigned",
                               "openid4vp-v1-signed",
                               "openid4vp-v1-multisigned",
                               "org-iso-mdoc",
                               "openid4vci",
                               "openid4vci-v1"};
    for (const String& protocol : kAllowed) {
      EXPECT_TRUE(
          DigitalCredential::userAgentAllowsProtocol(script_state, protocol))
          << "Expected valid protocol was rejected: " << protocol;
    }

    const String kBlocked[] = {"openid4vp", "some-random-protocol",
                               "INVALID_PROTOCOL"};
    for (const String& protocol : kBlocked) {
      EXPECT_FALSE(
          DigitalCredential::userAgentAllowsProtocol(script_state, protocol))
          << "Expected invalid protocol was allowed: " << protocol;
    }
  }
}

}  // namespace blink
