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

#include "components/media_router/common/providers/cast/channel/cast_auth_util.h"

#include <cstdlib>
#include <iostream>
#include <string>
#include <vector>

#include "base/no_destructor.h"
#include "base/notreached.h"
#include "base/time/time_override.h"
#include "components/media_router/common/providers/cast/certificate/cast_cert_reader.h"
#include "components/media_router/common/providers/cast/certificate/cast_cert_test_helpers.h"
#include "components/media_router/common/providers/cast/channel/cast_auth_util_fuzzer_shared.h"
#include "components/media_router/common/providers/cast/channel/fuzz_proto/fuzzer_inputs.pb.h"
#include "components/media_router/common/providers/cast/channel/fuzz_proto/fuzzer_inputs_fuzzable.pb.h"
// Generated by the "cast_auth_util_fuzzer_certs" data_headers target.
#include "components/test/data/media_router/common/providers/cast/certificate/certificates/chromecast_gen1_data.h"
#include "net/cert/x509_certificate.h"
#include "net/cert/x509_util.h"
#include "net/test/test_certificate_data.h"
#include "testing/libfuzzer/proto/lpm_interface.h"

namespace cast_channel {
namespace fuzz {
namespace {

const uint8_t kCertData[] = {
// Generated by //net/data/ssl/certificates:generate_fuzzer_cert_includes
#include "net/data/ssl/certificates/wildcard.inc"
};

base::NoDestructor<std::vector<std::string>> certs;

static bool InitializeOnce() {
  *certs = cast_certificate::ReadCertificateChainFromString(
      openscreen::cast::kChromecastGen1);
  CHECK(certs->size() >= 1)
      << "We should always have at least one certificate.";
  return true;
}

base::Time UpdateTime(TimeBoundCase c, int direction) {
  switch (c) {
    case TimeBoundCase::VALID:
      // Create bound that include the current date.
      return base::Time::Now() + base::Days(direction);
    case TimeBoundCase::INVALID:
      // Create a bound that excludes the current date.
      return base::Time::Now() + base::Days(-direction);
    case TimeBoundCase::OOB:
      // Create a bound so far in the past/future it's not valid.
      return base::Time::Now() + base::Days(direction * 10000);
    case TimeBoundCase::MISSING:
      // Remove any existing bound.
      return base::Time();
    default:
      NOTREACHED();
  }
}

DEFINE_PROTO_FUZZER(const fuzzable::cast_channel::fuzz::CastAuthUtilInputs&
                        fuzzable_input_union) {
  static bool init = InitializeOnce();
  CHECK(init);

  std::string serialized;
  CHECK(fuzzable_input_union.SerializeToString(&serialized));
  CastAuthUtilInputs input_union;
  // Recursion limits can cause parsing to fail.
  if (!input_union.ParseFromString(serialized)) {
    return;
  }

  if (input_union.input_case() !=
      CastAuthUtilInputs::kAuthenticateChallengeReplyInput) {
    return;
  }

  auto& input = *input_union.mutable_authenticate_challenge_reply_input();

  SetupAuthenticateChallengeReplyInput(*certs, &input);

  // Build a well-formed cert with start and expiry times relative to the
  // current time.  The actual cert doesn't matter for testing purposes
  // because validation failures are ignored.
  scoped_refptr<net::X509Certificate> peer_cert =
      net::X509Certificate::CreateFromBytes(kCertData);
  peer_cert->set_valid_start_for_testing(UpdateTime(input.start_case(), -1));
  peer_cert->set_valid_expiry_for_testing(UpdateTime(input.expiry_case(), +1));

  AuthContext context = AuthContext::CreateForTest(input.nonce());

  AuthenticateChallengeReply(input.cast_message(), *peer_cert, context);
}

}  // namespace
}  // namespace fuzz
}  // namespace cast_channel
