/*
 *  Copyright (c) 2015 The WebRTC project authors. All Rights Reserved.
 *
 *  Use of this source code is governed by a BSD-style license
 *  that can be found in the LICENSE file in the root of the source
 *  tree. An additional intellectual property rights grant can be found
 *  in the file PATENTS.  All contributing project authors may
 *  be found in the AUTHORS file in the root of the source tree.
 */

#include "media/engine/webrtc_media_engine.h"

#include <string>
#include <vector>

#include "absl/strings/string_view.h"
#include "api/field_trials.h"
#include "api/rtp_header_extension_id.h"
#include "api/rtp_parameters.h"
#include "test/create_test_field_trials.h"
#include "test/gtest.h"

namespace webrtc {
namespace {

std::vector<RtpExtension> MakeUniqueExtensions() {
  std::vector<RtpExtension> result;
  char name[] = "a";
  for (int i = 0; i < 7; ++i) {
    result.push_back(RtpExtension(name, RtpHeaderExtensionId(1 + i)));
    name[0]++;
    result.push_back(RtpExtension(name, RtpHeaderExtensionId(255 - i)));
    name[0]++;
  }
  return result;
}

std::vector<RtpExtension> MakeRedundantExtensions() {
  std::vector<RtpExtension> result;
  char name[] = "a";
  for (int i = 0; i < 7; ++i) {
    result.push_back(RtpExtension(name, RtpHeaderExtensionId(1 + i)));
    result.push_back(RtpExtension(name, RtpHeaderExtensionId(255 - i)));
    name[0]++;
  }
  return result;
}

bool SupportedExtensions1(absl::string_view name) {
  return name == "c" || name == "i";
}

bool SupportedExtensions2(absl::string_view name) {
  return name != "a" && name != "n";
}

bool IsSorted(const std::vector<RtpExtension>& extensions) {
  const std::string* last = nullptr;
  for (const auto& extension : extensions) {
    if (last && *last > extension.uri) {
      return false;
    }
    last = &extension.uri;
  }
  return true;
}
}  // namespace

TEST(WebRtcMediaEngineTest, FilterRtpExtensionsEmptyList) {
  std::vector<RtpExtension> extensions;
  FieldTrials trials = CreateTestFieldTrials();
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions1, true, trials);
  EXPECT_EQ(0u, filtered.size());
}

TEST(WebRtcMediaEngineTest, FilterRtpExtensionsIncludeOnlySupported) {
  std::vector<RtpExtension> extensions = MakeUniqueExtensions();
  FieldTrials trials = CreateTestFieldTrials();
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions1, false, trials);
  EXPECT_EQ(2u, filtered.size());
  EXPECT_EQ("c", filtered[0].uri);
  EXPECT_EQ("i", filtered[1].uri);
}

TEST(WebRtcMediaEngineTest, FilterRtpExtensionsSortedByName1) {
  std::vector<RtpExtension> extensions = MakeUniqueExtensions();
  FieldTrials trials = CreateTestFieldTrials();
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions2, false, trials);
  EXPECT_EQ(12u, filtered.size());
  EXPECT_TRUE(IsSorted(filtered));
}

TEST(WebRtcMediaEngineTest, FilterRtpExtensionsSortedByName2) {
  std::vector<RtpExtension> extensions = MakeUniqueExtensions();
  FieldTrials trials = CreateTestFieldTrials();
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions2, true, trials);
  EXPECT_EQ(12u, filtered.size());
  EXPECT_TRUE(IsSorted(filtered));
}

TEST(WebRtcMediaEngineTest, FilterRtpExtensionsDontRemoveRedundant) {
  std::vector<RtpExtension> extensions = MakeRedundantExtensions();
  FieldTrials trials = CreateTestFieldTrials();
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions2, false, trials);
  EXPECT_EQ(12u, filtered.size());
  EXPECT_TRUE(IsSorted(filtered));
  EXPECT_EQ(filtered[0].uri, filtered[1].uri);
}

TEST(WebRtcMediaEngineTest, FilterRtpExtensionsRemoveRedundant) {
  std::vector<RtpExtension> extensions = MakeRedundantExtensions();
  FieldTrials trials = CreateTestFieldTrials();
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions2, true, trials);
  EXPECT_EQ(6u, filtered.size());
  EXPECT_TRUE(IsSorted(filtered));
  EXPECT_NE(filtered[0].uri, filtered[1].uri);
}

TEST(WebRtcMediaEngineTest, FilterRtpExtensionsRemoveRedundantEncrypted1) {
  std::vector<RtpExtension> extensions;
  extensions.push_back(RtpExtension("b", RtpHeaderExtensionId(1)));
  extensions.push_back(RtpExtension("b", RtpHeaderExtensionId(2), true));
  extensions.push_back(RtpExtension("c", RtpHeaderExtensionId(3)));
  extensions.push_back(RtpExtension("b", RtpHeaderExtensionId(4)));
  FieldTrials trials = CreateTestFieldTrials();
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions2, true, trials);
  EXPECT_EQ(3u, filtered.size());
  EXPECT_TRUE(IsSorted(filtered));
  EXPECT_EQ(filtered[0].uri, filtered[1].uri);
  EXPECT_NE(filtered[0].encrypt, filtered[1].encrypt);
  EXPECT_NE(filtered[0].uri, filtered[2].uri);
  EXPECT_NE(filtered[1].uri, filtered[2].uri);
}

TEST(WebRtcMediaEngineTest, FilterRtpExtensionsRemoveRedundantEncrypted2) {
  std::vector<RtpExtension> extensions;
  extensions.push_back(RtpExtension("b", RtpHeaderExtensionId(1), true));
  extensions.push_back(RtpExtension("b", RtpHeaderExtensionId(2)));
  extensions.push_back(RtpExtension("c", RtpHeaderExtensionId(3)));
  extensions.push_back(RtpExtension("b", RtpHeaderExtensionId(4)));
  FieldTrials trials = CreateTestFieldTrials();
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions2, true, trials);
  EXPECT_EQ(3u, filtered.size());
  EXPECT_TRUE(IsSorted(filtered));
  EXPECT_EQ(filtered[0].uri, filtered[1].uri);
  EXPECT_NE(filtered[0].encrypt, filtered[1].encrypt);
  EXPECT_NE(filtered[0].uri, filtered[2].uri);
  EXPECT_NE(filtered[1].uri, filtered[2].uri);
}

TEST(WebRtcMediaEngineTest, FilterRtpExtensionsRemoveRedundantBwe1) {
  FieldTrials trials =
      CreateTestFieldTrials("WebRTC-FilterAbsSendTimeExtension/Enabled/");
  std::vector<RtpExtension> extensions;
  extensions.push_back(RtpExtension(RtpExtension::kTransportSequenceNumberUri,
                                    RtpHeaderExtensionId(3)));
  extensions.push_back(
      RtpExtension(RtpExtension::kTimestampOffsetUri, RtpHeaderExtensionId(9)));
  extensions.push_back(
      RtpExtension(RtpExtension::kAbsSendTimeUri, RtpHeaderExtensionId(6)));
  extensions.push_back(RtpExtension(RtpExtension::kTransportSequenceNumberUri,
                                    RtpHeaderExtensionId(1)));
  extensions.push_back(RtpExtension(RtpExtension::kTimestampOffsetUri,
                                    RtpHeaderExtensionId(14)));
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions2, true, trials);
  EXPECT_EQ(1u, filtered.size());
  EXPECT_EQ(RtpExtension::kTransportSequenceNumberUri, filtered[0].uri);
}

TEST(WebRtcMediaEngineTest,
     FilterRtpExtensionsRemoveRedundantBwe1KeepAbsSendTime) {
  std::vector<RtpExtension> extensions;
  extensions.push_back(RtpExtension(RtpExtension::kTransportSequenceNumberUri,
                                    RtpHeaderExtensionId(3)));
  extensions.push_back(
      RtpExtension(RtpExtension::kTimestampOffsetUri, RtpHeaderExtensionId(9)));
  extensions.push_back(
      RtpExtension(RtpExtension::kAbsSendTimeUri, RtpHeaderExtensionId(6)));
  extensions.push_back(RtpExtension(RtpExtension::kTransportSequenceNumberUri,
                                    RtpHeaderExtensionId(1)));
  extensions.push_back(RtpExtension(RtpExtension::kTimestampOffsetUri,
                                    RtpHeaderExtensionId(14)));
  FieldTrials trials = CreateTestFieldTrials();
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions2, true, trials);
  EXPECT_EQ(2u, filtered.size());
  EXPECT_EQ(RtpExtension::kTransportSequenceNumberUri, filtered[0].uri);
  EXPECT_EQ(RtpExtension::kAbsSendTimeUri, filtered[1].uri);
}

TEST(WebRtcMediaEngineTest, FilterRtpExtensionsRemoveRedundantBweEncrypted1) {
  FieldTrials trials =
      CreateTestFieldTrials("WebRTC-FilterAbsSendTimeExtension/Enabled/");
  std::vector<RtpExtension> extensions;
  extensions.push_back(RtpExtension(RtpExtension::kTransportSequenceNumberUri,
                                    RtpHeaderExtensionId(3)));
  extensions.push_back(RtpExtension(RtpExtension::kTransportSequenceNumberUri,
                                    RtpHeaderExtensionId(4), true));
  extensions.push_back(
      RtpExtension(RtpExtension::kTimestampOffsetUri, RtpHeaderExtensionId(9)));
  extensions.push_back(
      RtpExtension(RtpExtension::kAbsSendTimeUri, RtpHeaderExtensionId(6)));
  extensions.push_back(RtpExtension(RtpExtension::kTransportSequenceNumberUri,
                                    RtpHeaderExtensionId(1)));
  extensions.push_back(RtpExtension(RtpExtension::kTransportSequenceNumberUri,
                                    RtpHeaderExtensionId(2), true));
  extensions.push_back(RtpExtension(RtpExtension::kTimestampOffsetUri,
                                    RtpHeaderExtensionId(14)));
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions2, true, trials);
  EXPECT_EQ(2u, filtered.size());
  EXPECT_EQ(RtpExtension::kTransportSequenceNumberUri, filtered[0].uri);
  EXPECT_EQ(RtpExtension::kTransportSequenceNumberUri, filtered[1].uri);
  EXPECT_NE(filtered[0].encrypt, filtered[1].encrypt);
}

TEST(WebRtcMediaEngineTest,
     FilterRtpExtensionsRemoveRedundantBweEncrypted1KeepAbsSendTime) {
  std::vector<RtpExtension> extensions;
  extensions.push_back(RtpExtension(RtpExtension::kTransportSequenceNumberUri,
                                    RtpHeaderExtensionId(3)));
  extensions.push_back(RtpExtension(RtpExtension::kTransportSequenceNumberUri,
                                    RtpHeaderExtensionId(4), true));
  extensions.push_back(
      RtpExtension(RtpExtension::kTimestampOffsetUri, RtpHeaderExtensionId(9)));
  extensions.push_back(
      RtpExtension(RtpExtension::kAbsSendTimeUri, RtpHeaderExtensionId(6)));
  extensions.push_back(RtpExtension(RtpExtension::kTransportSequenceNumberUri,
                                    RtpHeaderExtensionId(1)));
  extensions.push_back(RtpExtension(RtpExtension::kTransportSequenceNumberUri,
                                    RtpHeaderExtensionId(2), true));
  extensions.push_back(RtpExtension(RtpExtension::kTimestampOffsetUri,
                                    RtpHeaderExtensionId(14)));
  FieldTrials trials = CreateTestFieldTrials();
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions2, true, trials);
  EXPECT_EQ(3u, filtered.size());
  EXPECT_EQ(RtpExtension::kTransportSequenceNumberUri, filtered[0].uri);
  EXPECT_EQ(RtpExtension::kTransportSequenceNumberUri, filtered[1].uri);
  EXPECT_EQ(RtpExtension::kAbsSendTimeUri, filtered[2].uri);
  EXPECT_NE(filtered[0].encrypt, filtered[1].encrypt);
}

TEST(WebRtcMediaEngineTest, FilterRtpExtensionsRemoveRedundantBwe2) {
  std::vector<RtpExtension> extensions;
  extensions.push_back(
      RtpExtension(RtpExtension::kTimestampOffsetUri, RtpHeaderExtensionId(1)));
  extensions.push_back(
      RtpExtension(RtpExtension::kAbsSendTimeUri, RtpHeaderExtensionId(14)));
  extensions.push_back(
      RtpExtension(RtpExtension::kTimestampOffsetUri, RtpHeaderExtensionId(7)));
  FieldTrials trials = CreateTestFieldTrials();
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions2, true, trials);
  EXPECT_EQ(1u, filtered.size());
  EXPECT_EQ(RtpExtension::kAbsSendTimeUri, filtered[0].uri);
}

TEST(WebRtcMediaEngineTest, FilterRtpExtensionsRemoveRedundantBwe3) {
  std::vector<RtpExtension> extensions;
  extensions.push_back(
      RtpExtension(RtpExtension::kTimestampOffsetUri, RtpHeaderExtensionId(2)));
  extensions.push_back(RtpExtension(RtpExtension::kTimestampOffsetUri,
                                    RtpHeaderExtensionId(14)));
  FieldTrials trials = CreateTestFieldTrials();
  std::vector<RtpExtension> filtered =
      FilterRtpExtensions(extensions, SupportedExtensions2, true, trials);
  EXPECT_EQ(1u, filtered.size());
  EXPECT_EQ(RtpExtension::kTimestampOffsetUri, filtered[0].uri);
}

}  // namespace webrtc
