// Copyright 2012 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/dns/address_sorter_posix.h"

#include <memory>
#include <optional>
#include <string>
#include <string_view>
#include <vector>

#include "base/check_op.h"
#include "base/containers/span.h"
#include "base/functional/bind.h"
#include "base/memory/raw_ptr.h"
#include "base/notimplemented.h"
#include "base/notreached.h"
#include "base/strings/strcat.h"
#include "base/task/single_thread_task_runner.h"
#include "base/test/bind.h"
#include "base/test/run_until.h"
#include "base/test/scoped_feature_list.h"
#include "base/test/task_environment.h"
#include "net/base/features.h"
#include "net/base/ip_address.h"
#include "net/base/ip_endpoint.h"
#include "net/base/net_errors.h"
#include "net/base/network_anonymization_key.h"
#include "net/base/network_change_notifier.h"
#include "net/base/schemeful_site.h"
#include "net/base/test_completion_callback.h"
#include "net/log/net_log_with_source.h"
#include "net/socket/client_socket_factory.h"
#include "net/socket/datagram_client_socket.h"
#include "net/socket/socket_performance_watcher.h"
#include "net/socket/ssl_client_socket.h"
#include "net/socket/stream_socket.h"
#include "net/traffic_annotation/network_traffic_annotation.h"
#include "testing/gmock/include/gmock/gmock.h"
#include "testing/gtest/include/gtest/gtest.h"
#include "url/gurl.h"

namespace net {
namespace {

// Used to map destination address to source address.
typedef std::map<IPAddress, IPAddress> AddressMapping;

IPAddress ParseIP(std::string_view str) {
  IPAddress addr;
  CHECK(addr.AssignFromIPLiteral(str));
  return addr;
}

// A mock socket which binds to source address according to AddressMapping.
class TestUDPClientSocket : public DatagramClientSocket {
 public:
  enum class ConnectMode { kSynchronous, kAsynchronous, kAsynchronousManual };
  TestUDPClientSocket(const AddressMapping* mapping,
                      ConnectMode connect_mode,
                      handles::NetworkHandle target_network)
      : mapping_(mapping),
        connect_mode_(connect_mode),
        target_network_(target_network) {}

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

  ~TestUDPClientSocket() override = default;

  int Read(IOBuffer*, int, CompletionOnceCallback) override {
    NOTIMPLEMENTED();
    return OK;
  }
  base::expected<DatagramsMetadata, Error> ReadMultiple(
      IOBuffer* buf,
      size_t buf_len,
      size_t maximum_packet_size,
      base::OnceCallback<void(base::expected<DatagramsMetadata, Error>)>
          callback) override {
    NOTIMPLEMENTED();
    return base::unexpected(ERR_NOT_IMPLEMENTED);
  }

  int Write(IOBuffer*,
            int,
            CompletionOnceCallback,
            const NetworkTrafficAnnotationTag& traffic_annotation) override {
    NOTIMPLEMENTED();
    return OK;
  }
  int SetReceiveBufferSize(int32_t) override { return OK; }
  int SetSendBufferSize(int32_t) override { return OK; }
  int SetDoNotFragment() override { return OK; }
  int SetRecvTos() override { return OK; }
  int SetTos(DiffServCodePoint dscp, EcnCodePoint ecn) override { return OK; }

  void Close() override {}
  int GetPeerAddress(IPEndPoint* address) const override {
    NOTIMPLEMENTED();
    return OK;
  }
  int GetLocalAddress(IPEndPoint* address) const override {
    if (!connected_)
      return ERR_UNEXPECTED;
    *address = local_endpoint_;
    return OK;
  }
  void UseNonBlockingIO() override {}
  int SetMulticastInterface(uint32_t interface_index) override {
    NOTIMPLEMENTED();
    return ERR_NOT_IMPLEMENTED;
  }

  int ConnectUsingNetwork(handles::NetworkHandle network,
                          const IPEndPoint& address) override {
    NOTIMPLEMENTED();
    return ERR_NOT_IMPLEMENTED;
  }

  int ConnectUsingDefaultNetwork(const IPEndPoint& address) override {
    NOTIMPLEMENTED();
    return ERR_NOT_IMPLEMENTED;
  }

  int ConnectAsync(const IPEndPoint& address,
                   CompletionOnceCallback callback) override {
    DCHECK(callback);
    int rv = Connect(address);
    finish_connect_callback_ =
        base::BindOnce(&TestUDPClientSocket::RunConnectCallback,
                       weak_ptr_factory_.GetWeakPtr(), std::move(callback), rv);
    if (connect_mode_ == ConnectMode::kAsynchronous) {
      base::SingleThreadTaskRunner::GetCurrentDefault()->PostTask(
          FROM_HERE, std::move(finish_connect_callback_));
      return ERR_IO_PENDING;
    } else if (connect_mode_ == ConnectMode::kAsynchronousManual) {
      return ERR_IO_PENDING;
    }
    return rv;
  }

  int ConnectUsingNetworkAsync(handles::NetworkHandle network,
                               const IPEndPoint& address,
                               CompletionOnceCallback callback) override {
    NOTIMPLEMENTED();
    return ERR_NOT_IMPLEMENTED;
  }

  int ConnectUsingDefaultNetworkAsync(
      const IPEndPoint& address,
      CompletionOnceCallback callback) override {
    NOTIMPLEMENTED();
    return ERR_NOT_IMPLEMENTED;
  }

  handles::NetworkHandle GetBoundNetwork() const override {
    return target_network_;
  }
  void ApplySocketTag(const SocketTag& tag) override {}
  void SetMsgConfirm(bool confirm) override {}

  int Connect(const IPEndPoint& remote) override {
    if (connected_)
      return ERR_UNEXPECTED;
    auto it = mapping_->find(remote.address());
    if (it == mapping_->end())
      return ERR_FAILED;
    connected_ = true;
    local_endpoint_ = IPEndPoint(it->second, 39874 /* arbitrary port */);
    return OK;
  }

  const NetLogWithSource& NetLog() const override { return net_log_; }

  void FinishConnect() { std::move(finish_connect_callback_).Run(); }

  DscpAndEcn GetLastTos() const override { return {DSCP_DEFAULT, ECN_DEFAULT}; }

 private:
  void RunConnectCallback(CompletionOnceCallback callback, int rv) {
    std::move(callback).Run(rv);
  }
  NetLogWithSource net_log_;
  raw_ptr<const AddressMapping> mapping_;
  bool connected_ = false;
  IPEndPoint local_endpoint_;
  ConnectMode connect_mode_;
  base::OnceClosure finish_connect_callback_;
  handles::NetworkHandle target_network_ = handles::kInvalidNetworkHandle;

  base::WeakPtrFactory<TestUDPClientSocket> weak_ptr_factory_{this};
};

// Creates TestUDPClientSockets and maintains an AddressMapping.
class TestSocketFactory : public ClientSocketFactory {
 public:
  TestSocketFactory() = default;

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

  ~TestSocketFactory() override = default;

  std::unique_ptr<DatagramClientSocket> CreateDatagramClientSocket(
      DatagramSocket::BindType,
      handles::NetworkHandle target_network,
      NetLog*,
      const NetLogSource&) override {
    // This is used only for testing in scenarios that do not involve multiple
    // networks. With that in mind, it's safe to ignore the `target_network`.
    auto new_socket = std::make_unique<TestUDPClientSocket>(
        &mapping_, connect_mode_, target_network);
    if (socket_create_callback_) {
      socket_create_callback_.Run(new_socket.get());
    }
    return new_socket;
  }
  std::unique_ptr<TransportClientSocket> CreateTransportClientSocket(
      const AddressList&,
      handles::NetworkHandle target_network,
      std::unique_ptr<SocketPerformanceWatcher>,
      net::NetworkQualityEstimator*,
      NetLog*,
      const NetLogSource&) override {
    // This is used only for testing in scenarios that do not involve multiple
    // networks. With that in mind, it's safe to ignore the `target_network`.
    NOTIMPLEMENTED();
    return nullptr;
  }
  std::unique_ptr<SSLClientSocket> CreateSSLClientSocket(
      SSLClientContext*,
      std::unique_ptr<StreamSocket>,
      const HostPortPair&,
      const SSLConfig&) override {
    NOTIMPLEMENTED();
    return nullptr;
  }
  void AddMapping(const IPAddress& dst, const IPAddress& src) {
    mapping_[dst] = src;
  }
  void SetConnectMode(TestUDPClientSocket::ConnectMode connect_mode) {
    connect_mode_ = connect_mode;
  }
  void SetSocketCreateCallback(
      base::RepeatingCallback<void(TestUDPClientSocket*)>
          socket_create_callback) {
    socket_create_callback_ = std::move(socket_create_callback);
  }

 private:
  AddressMapping mapping_;
  TestUDPClientSocket::ConnectMode connect_mode_ =
      TestUDPClientSocket::ConnectMode::kSynchronous;
  base::RepeatingCallback<void(TestUDPClientSocket*)> socket_create_callback_;
};

void OnSortComplete(bool& completed,
                    std::vector<IPEndPoint>* sorted_buf,
                    CompletionOnceCallback callback,
                    bool success,
                    std::vector<IPEndPoint> sorted) {
  EXPECT_TRUE(success);
  completed = true;
  if (success)
    *sorted_buf = std::move(sorted);
  std::move(callback).Run(OK);
}

}  // namespace

// TaskEnvironment is required to register an IPAddressObserver from the
// constructor of AddressSorterPosix.
class AddressSorterPosixTest : public ::testing::Test {
 protected:
  void SetUp() override {
    task_environment_.emplace();
    sorter_ = std::make_unique<AddressSorterPosix>(&socket_factory_);
  }

  void TearDown() override {
    sorter_.reset();
    task_environment_.reset();
  }

  void AddMapping(const std::string& dst, const std::string& src) {
    socket_factory_.AddMapping(ParseIP(dst), ParseIP(src));
  }

  void SetSocketCreateCallback(
      base::RepeatingCallback<void(TestUDPClientSocket*)>
          socket_create_callback) {
    socket_factory_.SetSocketCreateCallback(std::move(socket_create_callback));
  }

  void SetConnectMode(TestUDPClientSocket::ConnectMode connect_mode) {
    socket_factory_.SetConnectMode(connect_mode);
  }

  AddressSorterPosix::SourceAddressInfo* GetSourceInfo(
      const std::string& addr) {
    IPAddress address = ParseIP(addr);
    AddressSorterPosix::SourceAddressInfo* info =
        &sorter_->source_map_[address];
    if (info->scope == AddressSorterPosix::SCOPE_UNDEFINED)
      sorter_->FillPolicy(address, info);
    return info;
  }

  void RunUntilNotificationsDelivered() {
    // This will make the presubmit complain, but there doesn't seem to be a
    // better way of doing this.
    task_environment_->RunUntilIdle();
  }

  TestSocketFactory socket_factory_;
  std::unique_ptr<AddressSorterPosix> sorter_;
  bool completed_ = false;

 private:
  friend class AddressSorterPosixSyncOrAsyncTest;

  // This is wrapped in std::optional so that it can be constructed after the
  // ScopedFeatureList objects used by subclasses and destroyed before them.
  std::optional<base::test::TaskEnvironment> task_environment_;
};

// Parameterized subclass of AddressSorterPosixTest. Necessary because not every
// test needs to be parameterized.
class AddressSorterPosixSyncOrAsyncTest
    : public AddressSorterPosixTest,
      public testing::WithParamInterface<
          std::tuple<TestUDPClientSocket::ConnectMode, bool>> {
 protected:
  AddressSorterPosixSyncOrAsyncTest() {
    SetConnectMode(std::get<0>(GetParam()));
    feature_list_.InitWithFeatureState(features::kAddressSorterConnectCache,
                                       std::get<1>(GetParam()));
  }

  // Verify |addresses| matches |order| after sorting.
  void Verify(base::span<const std::string_view> addresses,
              base::span<const int> order) {
    std::vector<IPEndPoint> endpoints;
    for (auto addr : addresses) {
      endpoints.emplace_back(ParseIP(addr), 80);
    }
    for (auto order_i : order) {
      CHECK_LT(order_i, static_cast<int>(endpoints.size()));
    }

    std::vector<IPEndPoint> sorted;
    TestCompletionCallback callback;
    sorter_->Sort(endpoints, NetworkAnonymizationKey(),
                  // This is used only for testing in scenarios that do not
                  // involve multiple networks. With that in mind, it's safe to
                  // always use the default network.
                  handles::kInvalidNetworkHandle,
                  base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                                 callback.callback()));
    callback.WaitForResult();

    for (size_t i = 0; (i < sorted.size()) || (i < order.size()); ++i) {
      IPEndPoint expected =
          i < order.size() ? endpoints[order[i]] : IPEndPoint();
      IPEndPoint actual = i < sorted.size() ? sorted[i] : IPEndPoint();
      EXPECT_TRUE(expected == actual)
          << "Endpoint out of order at position " << i << "\n"
          << "  Actual: " << actual.ToString() << "\n"
          << "Expected: " << expected.ToString();
    }
    EXPECT_TRUE(completed_);
  }

 private:
  base::test::ScopedFeatureList feature_list_;
};

INSTANTIATE_TEST_SUITE_P(
    AddressSorterPosix,
    AddressSorterPosixSyncOrAsyncTest,
    ::testing::Combine(
        ::testing::Values(TestUDPClientSocket::ConnectMode::kSynchronous,
                          TestUDPClientSocket::ConnectMode::kAsynchronous),
        ::testing::Bool()),
    ([](const ::testing::TestParamInfo<
         AddressSorterPosixSyncOrAsyncTest::ParamType>& info) {
      const auto& [connect_mode, enable_cache] = info.param;
      return base::StrCat({
          "ConnectMode_",
          connect_mode == TestUDPClientSocket::ConnectMode::kSynchronous
              ? "Sync"
              : "Async",
          "_Cache_",
          enable_cache ? "Enabled" : "Disabled",
      });
    }));

// Rule 1: Avoid unusable destinations.
TEST_P(AddressSorterPosixSyncOrAsyncTest, Rule1) {
  AddMapping("10.0.0.231", "10.0.0.1");
  const std::string_view addresses[] = {"::1", "10.0.0.231", "127.0.0.1"};
  const int order[] = {1};
  Verify(addresses, order);
}

// Rule 2: Prefer matching scope.
TEST_P(AddressSorterPosixSyncOrAsyncTest, Rule2) {
  AddMapping("3002:01::119", "4000::10");  // matching global
  AddMapping("ff32::1", "fe81::10");      // matching link-local
  AddMapping("fec1:01::119", "fec1:01::10");  // matching node-local
  AddMapping("3002:02::119", "::1");          // global vs. link-local
  AddMapping("fec1:02::119", "fe81::10");     // site-local vs. link-local
  AddMapping("8.0.0.1", "169.254.0.10");  // global vs. link-local
  // In all three cases, matching scope is preferred.
  const int order[] = {1, 0};
  const std::string_view addresses1[] = {"3002:02::119", "3002:01::119"};
  Verify(addresses1, order);
  const std::string_view addresses2[] = {"fec1:02::119", "ff32::1"};
  Verify(addresses2, order);
  const std::string_view addresses3[] = {"8.0.0.1", "fec1:01::119"};
  Verify(addresses3, order);
}

// Rule 3: Avoid deprecated addresses.
TEST_P(AddressSorterPosixSyncOrAsyncTest, Rule3) {
  // Matching scope.
  AddMapping("3002:01::119", "4000::10");
  GetSourceInfo("4000::10")->deprecated = true;
  AddMapping("3002:02::119", "4000::20");
  const std::string_view addresses[] = {"3002:01::119", "3002:02::119"};
  const int order[] = {1, 0};
  Verify(addresses, order);
}

// Rule 4: Prefer home addresses.
TEST_P(AddressSorterPosixSyncOrAsyncTest, Rule4) {
  AddMapping("3002:01::119", "4000::10");
  AddMapping("3002:02::119", "4000::20");
  GetSourceInfo("4000::20")->home = true;
  const std::string_view addresses[] = {"3002:01::119", "3002:02::119"};
  const int order[] = {1, 0};
  Verify(addresses, order);
}

// Rule 5: Prefer matching label.
TEST_P(AddressSorterPosixSyncOrAsyncTest, Rule5) {
  AddMapping("::1", "::1");                       // matching loopback
  AddMapping("::ffff:1234:1", "::ffff:1234:10");  // matching IPv4-mapped
  AddMapping("2001::1", "::ffff:1234:10");        // Teredo vs. IPv4-mapped
  AddMapping("2002::1", "2001::10");              // 6to4 vs. Teredo
  const int order[] = {1, 0};
  {
    const std::string_view addresses[] = {"2001::1", "::1"};
    Verify(addresses, order);
  }
  {
    const std::string_view addresses[] = {"2002::1", "::ffff:1234:1"};
    Verify(addresses, order);
  }
}

// Rule 6: Prefer higher precedence.
TEST_P(AddressSorterPosixSyncOrAsyncTest, Rule6) {
  AddMapping("::1", "::1");                       // loopback
  AddMapping("ff32::1", "fe81::10");              // multicast
  AddMapping("::ffff:1234:1", "::ffff:1234:10");  // IPv4-mapped
  AddMapping("2001::1", "2001::10");              // Teredo
  const std::string_view addresses[] = {"2001::1", "::ffff:1234:1", "ff32::1",
                                        "::1"};
  const int order[] = {3, 2, 1, 0};
  Verify(addresses, order);
}

// Rule 7: Prefer native transport.
TEST_P(AddressSorterPosixSyncOrAsyncTest, Rule7) {
  AddMapping("3002:01::119", "4000::10");
  AddMapping("3002:02::119", "4000::20");
  GetSourceInfo("4000::20")->native = true;
  const std::string_view addresses[] = {"3002:01::119", "3002:02::119"};
  const int order[] = {1, 0};
  Verify(addresses, order);
}

// Rule 8: Prefer smaller scope.
TEST_P(AddressSorterPosixSyncOrAsyncTest, Rule8) {
  // Matching scope. Should precede the others by Rule 2.
  AddMapping("fe81::1", "fe81::10");  // link-local
  AddMapping("3000::1", "4000::10");  // global
  // Mismatched scope.
  AddMapping("ff32::1", "4000::10");  // link-local
  AddMapping("ff35::1", "4000::10");  // site-local
  AddMapping("ff38::1", "4000::10");  // org-local
  const std::string_view addresses[] = {"ff38::1", "3000::1", "ff35::1",
                                        "ff32::1", "fe81::1"};
  const int order[] = {4, 1, 3, 2, 0};
  Verify(addresses, order);
}

// Rule 9: Use longest matching prefix.
TEST_P(AddressSorterPosixSyncOrAsyncTest, Rule9) {
  AddMapping("3000::1", "3000:ffff::10");  // 16 bit match
  GetSourceInfo("3000:ffff::10")->prefix_length = 16;
  AddMapping("4000::1", "4000::10");       // 123 bit match, limited to 15
  GetSourceInfo("4000::10")->prefix_length = 15;
  AddMapping("4002::1", "4000::10");       // 14 bit match
  AddMapping("4080::1", "4000::10");       // 8 bit match
  const std::string_view addresses[] = {"4080::1", "4002::1", "4000::1",
                                        "3000::1"};
  const int order[] = {3, 2, 1, 0};
  Verify(addresses, order);
}

// Rule 10: Leave the order unchanged.
TEST_P(AddressSorterPosixSyncOrAsyncTest, Rule10) {
  AddMapping("4000::1", "4000::10");
  AddMapping("4000::2", "4000::10");
  AddMapping("4000::3", "4000::10");
  const std::string_view addresses[] = {"4000::1", "4000::2", "4000::3"};
  const int order[] = {0, 1, 2};
  Verify(addresses, order);
}

TEST_P(AddressSorterPosixSyncOrAsyncTest, MultipleRules) {
  AddMapping("::1", "::1");           // loopback
  AddMapping("ff32:01::119", "fe81::10");  // link-local multicast
  AddMapping("ff3e::1", "4000::10");  // global multicast
  AddMapping("4000::1", "4000::10");  // global unicast
  AddMapping("ff32:02::119", "fe81::20");  // deprecated link-local multicast
  GetSourceInfo("fe81::20")->deprecated = true;
  const std::string_view addresses[] = {
      "ff3e::1", "ff32:02::119", "4000::1", "ff32:01::119", "::1", "8.0.0.1"};
  const int order[] = {4, 3, 0, 2, 1};
  Verify(addresses, order);
}

TEST_P(AddressSorterPosixSyncOrAsyncTest, InputPortsAreMaintained) {
  AddMapping("::1", "::1");
  AddMapping("::2", "::2");
  AddMapping("::3", "::3");

  IPEndPoint endpoint1(ParseIP("::1"), /*port=*/111);
  IPEndPoint endpoint2(ParseIP("::2"), /*port=*/222);
  IPEndPoint endpoint3(ParseIP("::3"), /*port=*/333);

  std::vector<IPEndPoint> input = {endpoint1, endpoint2, endpoint3};
  std::vector<IPEndPoint> sorted;
  TestCompletionCallback callback;
  sorter_->Sort(input, NetworkAnonymizationKey(),
                handles::kInvalidNetworkHandle,
                base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                               callback.callback()));
  callback.WaitForResult();

  EXPECT_THAT(sorted, testing::ElementsAre(endpoint1, endpoint2, endpoint3));
}

TEST_P(AddressSorterPosixSyncOrAsyncTest, AddressSorterPosixDestroyed) {
  AddMapping("::1", "::1");
  AddMapping("::2", "::2");
  AddMapping("::3", "::3");

  IPEndPoint endpoint1(ParseIP("::1"), /*port=*/111);
  IPEndPoint endpoint2(ParseIP("::2"), /*port=*/222);
  IPEndPoint endpoint3(ParseIP("::3"), /*port=*/333);

  std::vector<IPEndPoint> input = {endpoint1, endpoint2, endpoint3};
  std::vector<IPEndPoint> sorted;
  TestCompletionCallback callback;
  sorter_->Sort(input, NetworkAnonymizationKey(),
                handles::kInvalidNetworkHandle,
                base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                               callback.callback()));
  sorter_.reset();
  base::RunLoop().RunUntilIdle();

  TestUDPClientSocket::ConnectMode connect_mode = std::get<0>(GetParam());
  if (connect_mode == TestUDPClientSocket::ConnectMode::kAsynchronous) {
    EXPECT_FALSE(completed_);
  } else {
    EXPECT_TRUE(completed_);
  }
}

TEST_F(AddressSorterPosixTest, RandomAsyncSocketOrder) {
  SetConnectMode(TestUDPClientSocket::ConnectMode::kAsynchronousManual);
  std::vector<TestUDPClientSocket*> created_sockets;
  SetSocketCreateCallback(base::BindRepeating(
      [](std::vector<TestUDPClientSocket*>& created_sockets,
         TestUDPClientSocket* socket) { created_sockets.push_back(socket); },
      std::ref(created_sockets)));

  AddMapping("::1", "::1");
  AddMapping("::2", "::2");
  AddMapping("::3", "::3");

  IPEndPoint endpoint1(ParseIP("::1"), /*port=*/111);
  IPEndPoint endpoint2(ParseIP("::2"), /*port=*/222);
  IPEndPoint endpoint3(ParseIP("::3"), /*port=*/333);

  std::vector<IPEndPoint> input = {endpoint1, endpoint2, endpoint3};
  std::vector<IPEndPoint> sorted;
  TestCompletionCallback callback;
  sorter_->Sort(input, NetworkAnonymizationKey(),
                handles::kInvalidNetworkHandle,
                base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                               callback.callback()));

  ASSERT_EQ(created_sockets.size(), 3u);
  created_sockets[1]->FinishConnect();
  created_sockets[2]->FinishConnect();
  created_sockets[0]->FinishConnect();

  base::RunLoop().RunUntilIdle();
  EXPECT_TRUE(completed_);
}

// Regression test for https://crbug.com/1374387
TEST_F(AddressSorterPosixTest, IPAddressChangedSort) {
  SetConnectMode(TestUDPClientSocket::ConnectMode::kAsynchronousManual);
  std::vector<TestUDPClientSocket*> created_sockets;
  SetSocketCreateCallback(base::BindRepeating(
      [](std::vector<TestUDPClientSocket*>& created_sockets,
         TestUDPClientSocket* socket) { created_sockets.push_back(socket); },
      std::ref(created_sockets)));

  AddMapping("::1", "::1");
  AddMapping("::2", "::2");
  AddMapping("::3", "::3");

  IPEndPoint endpoint1(ParseIP("::1"), /*port=*/111);
  IPEndPoint endpoint2(ParseIP("::2"), /*port=*/222);
  IPEndPoint endpoint3(ParseIP("::3"), /*port=*/333);

  std::vector<IPEndPoint> input = {endpoint1, endpoint2, endpoint3};
  std::vector<IPEndPoint> sorted;
  TestCompletionCallback callback;
  sorter_->Sort(input, NetworkAnonymizationKey(),
                handles::kInvalidNetworkHandle,
                base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                               callback.callback()));

  ASSERT_EQ(created_sockets.size(), 3u);
  created_sockets[0]->FinishConnect();
  // Trigger OnIPAddressChanged() to reset `source_map_`
  NetworkChangeNotifier::NotifyObserversOfIPAddressChangeForTests();
  base::RunLoop().RunUntilIdle();
  created_sockets[1]->FinishConnect();
  created_sockets[2]->FinishConnect();

  base::RunLoop().RunUntilIdle();
  EXPECT_TRUE(completed_);
}

class AddressSorterPosixCacheTest : public AddressSorterPosixTest {
 protected:
  AddressSorterPosixCacheTest() {
    feature_list_.InitAndEnableFeature(features::kAddressSorterConnectCache);
  }

 private:
  base::test::ScopedFeatureList feature_list_;
};

TEST_F(AddressSorterPosixCacheTest, CacheHitBypassesSocketCreation) {
  size_t socket_create_count = 0;
  SetSocketCreateCallback(base::BindLambdaForTesting(
      [&socket_create_count](TestUDPClientSocket* socket) {
        socket_create_count++;
      }));

  AddMapping("10.0.0.1", "10.0.0.10");
  AddMapping("8.8.8.8", "10.0.0.10");

  IPEndPoint endpoint1(ParseIP("10.0.0.1"), 80);
  IPEndPoint endpoint2(ParseIP("8.8.8.8"), 80);

  // First sort should miss cache and create sockets.
  {
    std::vector<IPEndPoint> sorted;
    TestCompletionCallback callback;
    sorter_->Sort({endpoint1, endpoint2}, NetworkAnonymizationKey(),
                  handles::kInvalidNetworkHandle,
                  base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                                 callback.callback()));
    callback.WaitForResult();
    EXPECT_EQ(socket_create_count, 2u);
  }

  // Second sort with same NAK should hit cache and bypass socket creation.
  {
    socket_create_count = 0;
    std::vector<IPEndPoint> sorted;
    TestCompletionCallback callback;
    sorter_->Sort({endpoint1, endpoint2}, NetworkAnonymizationKey(),
                  handles::kInvalidNetworkHandle,
                  base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                                 callback.callback()));
    callback.WaitForResult();
    EXPECT_EQ(socket_create_count, 0u);
  }
}

TEST_F(AddressSorterPosixCacheTest, StatePartitioningByNAK) {
  size_t socket_create_count = 0;
  SetSocketCreateCallback(base::BindLambdaForTesting(
      [&socket_create_count](TestUDPClientSocket* socket) {
        socket_create_count++;
      }));

  AddMapping("10.0.0.1", "10.0.0.10");
  IPEndPoint endpoint(ParseIP("10.0.0.1"), 80);

  SchemefulSite site_a(GURL("https://site_a.test/"));
  NetworkAnonymizationKey nak_a =
      NetworkAnonymizationKey::CreateSameSite(site_a);

  SchemefulSite site_b(GURL("https://site_b.test/"));
  NetworkAnonymizationKey nak_b =
      NetworkAnonymizationKey::CreateSameSite(site_b);

  {
    std::vector<IPEndPoint> sorted;
    TestCompletionCallback callback;
    sorter_->Sort({endpoint}, nak_a, handles::kInvalidNetworkHandle,
                  base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                                 callback.callback()));
    callback.WaitForResult();
    EXPECT_EQ(socket_create_count, 1u);
  }

  // Second sort with nak_b should miss cache.
  {
    socket_create_count = 0;
    std::vector<IPEndPoint> sorted;
    TestCompletionCallback callback;
    sorter_->Sort({endpoint}, nak_b, handles::kInvalidNetworkHandle,
                  base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                                 callback.callback()));
    callback.WaitForResult();
    EXPECT_EQ(socket_create_count, 1u);
  }
}

TEST_F(AddressSorterPosixCacheTest, StatePartitioningByTargetNetwork) {
  size_t socket_create_count = 0;
  handles::NetworkHandle last_bound_network = handles::kInvalidNetworkHandle;
  SetSocketCreateCallback(base::BindLambdaForTesting(
      [&socket_create_count, &last_bound_network](TestUDPClientSocket* socket) {
        socket_create_count++;
        last_bound_network = socket->GetBoundNetwork();
      }));

  AddMapping("10.0.0.1", "10.0.0.10");
  IPEndPoint endpoint(ParseIP("10.0.0.1"), 80);

  handles::NetworkHandle network_a = 1;
  handles::NetworkHandle network_b = 2;

  {
    std::vector<IPEndPoint> sorted;
    TestCompletionCallback callback;
    sorter_->Sort({endpoint}, NetworkAnonymizationKey(), network_a,
                  base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                                 callback.callback()));
    callback.WaitForResult();
    EXPECT_EQ(socket_create_count, 1u);
    EXPECT_EQ(last_bound_network, network_a);
  }

  // Second sort with network_b should miss cache.
  {
    socket_create_count = 0;
    last_bound_network = handles::kInvalidNetworkHandle;
    std::vector<IPEndPoint> sorted;
    TestCompletionCallback callback;
    sorter_->Sort({endpoint}, NetworkAnonymizationKey(), network_b,
                  base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                                 callback.callback()));
    callback.WaitForResult();
    EXPECT_EQ(socket_create_count, 1u);
    EXPECT_EQ(last_bound_network, network_b);
  }

  // Third sort with network_a should hit cache.
  {
    socket_create_count = 0;
    last_bound_network = handles::kInvalidNetworkHandle;
    std::vector<IPEndPoint> sorted;
    TestCompletionCallback callback;
    sorter_->Sort({endpoint}, NetworkAnonymizationKey(), network_a,
                  base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                                 callback.callback()));
    callback.WaitForResult();
    EXPECT_EQ(socket_create_count, 0u);
    EXPECT_EQ(last_bound_network, handles::kInvalidNetworkHandle);
  }
}

TEST_F(AddressSorterPosixCacheTest, SubnetMaskingMatchesSameCacheEntry) {
  size_t socket_create_count = 0;
  SetSocketCreateCallback(base::BindLambdaForTesting(
      [&socket_create_count](TestUDPClientSocket* socket) {
        socket_create_count++;
      }));

  AddMapping("10.0.0.1", "10.0.0.10");
  AddMapping("10.0.0.2", "10.0.0.10");
  AddMapping("2001:db8::1", "2001:db8::10");
  AddMapping("2001:db8::2", "2001:db8::10");

  IPEndPoint endpoint_v4_1(ParseIP("10.0.0.1"), 80);
  IPEndPoint endpoint_v4_2(ParseIP("10.0.0.2"), 80);
  IPEndPoint endpoint_v6_1(ParseIP("2001:db8::1"), 80);
  IPEndPoint endpoint_v6_2(ParseIP("2001:db8::2"), 80);

  {
    std::vector<IPEndPoint> sorted;
    TestCompletionCallback callback;
    sorter_->Sort({endpoint_v4_1, endpoint_v6_1}, NetworkAnonymizationKey(),
                  handles::kInvalidNetworkHandle,
                  base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                                 callback.callback()));
    callback.WaitForResult();
    EXPECT_EQ(socket_create_count, 2u);
  }

  // Second sort with different IPs in the same /24 and /64 should hit cache.
  {
    socket_create_count = 0;
    std::vector<IPEndPoint> sorted;
    TestCompletionCallback callback;
    sorter_->Sort({endpoint_v4_2, endpoint_v6_2}, NetworkAnonymizationKey(),
                  handles::kInvalidNetworkHandle,
                  base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                                 callback.callback()));
    callback.WaitForResult();
    EXPECT_EQ(socket_create_count, 0u);
  }
}

TEST_F(AddressSorterPosixCacheTest,
       SubnetMaskingDistinguishesDifferentSubnets) {
  size_t socket_create_count = 0;
  SetSocketCreateCallback(base::BindLambdaForTesting(
      [&socket_create_count](TestUDPClientSocket* socket) {
        socket_create_count++;
      }));

  // Third byte of IPv4 differs (10.0.0.1 vs 10.0.1.1).
  AddMapping("10.0.0.1", "10.0.0.10");
  AddMapping("10.0.1.1", "10.0.0.10");

  // 8th byte of IPv6 differs (2001:db8:0:0::1 vs 2001:db8:0:1::1).
  AddMapping("2001:db8::1", "2001:db8::10");
  AddMapping("2001:db8:0:1::1", "2001:db8::10");

  IPEndPoint endpoint_v4_1(ParseIP("10.0.0.1"), 80);
  IPEndPoint endpoint_v4_2(ParseIP("10.0.1.1"), 80);
  IPEndPoint endpoint_v6_1(ParseIP("2001:db8::1"), 80);
  IPEndPoint endpoint_v6_2(ParseIP("2001:db8:0:1::1"), 80);

  {
    std::vector<IPEndPoint> sorted;
    TestCompletionCallback callback;
    sorter_->Sort({endpoint_v4_1, endpoint_v6_1}, NetworkAnonymizationKey(),
                  handles::kInvalidNetworkHandle,
                  base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                                 callback.callback()));
    callback.WaitForResult();
    EXPECT_EQ(socket_create_count, 2u);
  }

  // Second sort with different subnets should miss cache and allocate sockets.
  {
    socket_create_count = 0;
    std::vector<IPEndPoint> sorted;
    TestCompletionCallback callback;
    sorter_->Sort({endpoint_v4_2, endpoint_v6_2}, NetworkAnonymizationKey(),
                  handles::kInvalidNetworkHandle,
                  base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                                 callback.callback()));
    callback.WaitForResult();
    EXPECT_EQ(socket_create_count, 2u);
  }
}

TEST_F(AddressSorterPosixCacheTest, NetworkChangeDuringSort) {
  SetConnectMode(TestUDPClientSocket::ConnectMode::kAsynchronousManual);
  std::vector<TestUDPClientSocket*> created_sockets;
  SetSocketCreateCallback(base::BindRepeating(
      [](std::vector<TestUDPClientSocket*>& created_sockets,
         TestUDPClientSocket* socket) { created_sockets.push_back(socket); },
      std::ref(created_sockets)));

  AddMapping("::1", "::1");
  AddMapping("::2", "::2");

  IPEndPoint endpoint1(ParseIP("::1"), /*port=*/111);
  IPEndPoint endpoint2(ParseIP("::2"), /*port=*/222);

  std::vector<IPEndPoint> input = {endpoint1, endpoint2};
  std::vector<IPEndPoint> sorted;
  TestCompletionCallback callback;
  sorter_->Sort(input, NetworkAnonymizationKey(),
                handles::kInvalidNetworkHandle,
                base::BindOnce(&OnSortComplete, std::ref(completed_), &sorted,
                               callback.callback()));

  ASSERT_EQ(created_sockets.size(), 2u);
  created_sockets[0]->FinishConnect();
  // Trigger OnNetworkChanged() to clear the cache.
  NetworkChangeNotifier::NotifyObserversOfNetworkChangeForTests(
      NetworkChangeNotifier::CONNECTION_UNKNOWN);
  RunUntilNotificationsDelivered();
  created_sockets[1]->FinishConnect();

  EXPECT_TRUE(base::test::RunUntil([&] { return completed_; }));
  // We should not have crashed.
}

}  // namespace net
