#include "common.hpp"
#include "handles/BackoffTimerHandleInterface.hpp"
#include "RTC/SCTP/association/AssociationListenerDeferrer.hpp"
#include "RTC/SCTP/association/StreamResetHandler.hpp"
#include "RTC/SCTP/packet/Packet.hpp"
#include "RTC/SCTP/packet/chunks/AnyForwardTsnChunk.hpp"
#include "RTC/SCTP/packet/chunks/ReConfigChunk.hpp"
#include "RTC/SCTP/packet/parameters/IncomingSsnResetRequestParameter.hpp"
#include "RTC/SCTP/packet/parameters/OutgoingSsnResetRequestParameter.hpp"
#include "RTC/SCTP/packet/parameters/ReconfigurationResponseParameter.hpp"
#include "RTC/SCTP/public/SctpOptions.hpp"
#include "RTC/SCTP/rx/DataTracker.hpp"
#include "RTC/SCTP/rx/ReassemblyQueue.hpp"
#include "RTC/SCTP/tx/RetransmissionQueue.hpp"
#include "RTC/SCTP/tx/RoundRobinSendQueue.hpp"
#include "mocks/include/MockShared.hpp"
#include "mocks/include/RTC/SCTP/association/MockAssociationListener.hpp"
#include "mocks/include/RTC/SCTP/association/MockTransmissionControlBlockContext.hpp"
#include <catch2/catch_test_macros.hpp>
#include <span>
#include <vector>

namespace
{
	constexpr uint32_t Arwnd{ 131072 };
	constexpr uint32_t LocalInitialTsn{ 0 };
	constexpr uint32_t RemoteInitialTsn{ 0 };
	constexpr int64_t InitialNowUs{ 10000 * 1000 };
	constexpr int64_t RtoMs{ 250 };

	/**
	 * A RTC::SCTP::StreamResetHandler under test, together with all the (real)
	 * components it depends on, its simulated clock and listener.
	 */
	class TestStreamResetHandler
	{
	public:
		TestStreamResetHandler()
		  // NOTE: The order in which these members are initialized is **critical**.
		  : shared(/*getTimeUs*/
			         [this]() -> int64_t
			         {
			           return this->nowUs;
		           }),
		    tcbContext(this->associationListener, this->sctpOptions),
		    delayedAckTimer(this->shared.CreateBackoffTimer(
		      BackoffTimerHandleInterface::BackoffTimerHandleOptions{
		        .listener         = std::addressof(this->backoffTimerListener),
		        .label            = "mock-sctp-delayed-ack",
		        .baseTimeoutMs    = this->sctpOptions.initialRtoMs,
		        .backoffAlgorithm = BackoffTimerHandleInterface::BackoffAlgorithm::FIXED })),
		    t3RtxTimer(this->shared.CreateBackoffTimer(
		      BackoffTimerHandleInterface::BackoffTimerHandleOptions{
		        .listener         = std::addressof(this->backoffTimerListener),
		        .label            = "mock-sctp-t3-rtx",
		        .baseTimeoutMs    = this->sctpOptions.initialRtoMs,
		        .backoffAlgorithm = BackoffTimerHandleInterface::BackoffAlgorithm::EXPONENTIAL })),
		    sendQueue(
		      this->associationListener,
		      this->sctpOptions.mtu,
		      this->sctpOptions.defaultStreamPriority,
		      this->sctpOptions.defaultStreamBufferedAmountLowThreshold,
		      this->sctpOptions.totalBufferedAmountLowThreshold),
		    dataTracker(this->delayedAckTimer.get(), RemoteInitialTsn),
		    reassemblyQueue(this->sctpOptions.maxReceiverWindowBufferSize),
		    retransmissionQueue(
		      std::addressof(this->retransmissionQueueListener),
		      this->associationListener,
		      LocalInitialTsn,
		      Arwnd,
		      this->sendQueue,
		      this->t3RtxTimer.get(),
		      this->sctpOptions,
		      /*supportsPartialReliability*/ true,
		      /*useMessageInterleaving*/ false),
		    streamResetHandler(
		      this->associationListenerDeferrer,
		      std::addressof(this->shared),
		      std::addressof(this->tcbContext),
		      std::addressof(this->dataTracker),
		      std::addressof(this->reassemblyQueue),
		      std::addressof(this->retransmissionQueue))
		{
			this->tcbContext.WillGetCurrentRtoUsOnce(
			  []()
			  {
				  return static_cast<int64_t>(RtoMs * 1000);
			  });
		}

	public:
		/**
		 * Advances the simulated clock by `incrementMs`.
		 *
		 * @remarks
		 * - The increment is given in milliseconds since it comes from the SCTP
		 *   options and the timers, which work in milliseconds.
		 */
		void AdvanceTimeMs(int64_t incrementMs)
		{
			this->nowUs += incrementMs * 1000;
		}

		/**
		 * Handles a received RE-CONFIG chunk within a deferrer scope, just like
		 * association does before dispatching to the handler.
		 */
		void HandleReceivedReConfigChunk(const RTC::SCTP::ReConfigChunk* receivedReConfigChunk)
		{
			const RTC::SCTP::AssociationListenerDeferrer::ScopedDeferrer deferrer(
			  this->associationListenerDeferrer);

			this->streamResetHandler.HandleReceivedReConfigChunk(receivedReConfigChunk);
		}

	private:
		/**
		 * A no-op BackoffTimerHandle listener for the timers owned directly by this
		 * fixture (the ones owned by the handler are driven by the handler itself).
		 */
		class BackoffTimerListener : public BackoffTimerHandleInterface::Listener
		{
		public:
			void OnBackoffTimer(
			  BackoffTimerHandleInterface* /*backoffTimer*/, int64_t& /*baseTimeoutMs*/, bool& /*stop*/) override
			{
			}
		};

		/**
		 * A no-op RTC::SCTP::RetransmissionQueue listener.
		 */
		class RetransmissionQueueListener : public RTC::SCTP::RetransmissionQueue::Listener
		{
		public:
			void OnRetransmissionQueueNewRttUs(int64_t /*rttUs*/) override
			{
			}

			void OnRetransmissionQueueClearRetransmissionCounter() override
			{
			}
		};

		// NOTE: Public members for testing.
	public:
		int64_t nowUs{ InitialNowUs };
		RTC::SCTP::SctpOptions sctpOptions;
		BackoffTimerListener backoffTimerListener;
		RetransmissionQueueListener retransmissionQueueListener;
		mocks::RTC::SCTP::MockAssociationListener associationListener;
		RTC::SCTP::AssociationListenerDeferrer associationListenerDeferrer{ std::addressof(
			this->associationListener) };
		mocks::MockShared shared;
		mocks::RTC::SCTP::MockTransmissionControlBlockContext tcbContext;
		const std::unique_ptr<BackoffTimerHandleInterface> delayedAckTimer;
		const std::unique_ptr<BackoffTimerHandleInterface> t3RtxTimer;
		RTC::SCTP::RoundRobinSendQueue sendQueue;
		RTC::SCTP::DataTracker dataTracker;
		RTC::SCTP::ReassemblyQueue reassemblyQueue;
		RTC::SCTP::RetransmissionQueue retransmissionQueue;
		RTC::SCTP::StreamResetHandler streamResetHandler;
	};
} // namespace

SCENARIO("SCTP RTC::SCTP::StreamResetHandler", "[sctp][streamresethandler]")
{
	SECTION("a chunk with no parameters returns an error")
	{
		TestStreamResetHandler test;

		std::vector<uint8_t> buffer(test.sctpOptions.mtu);

		const std::unique_ptr<RTC::SCTP::ReConfigChunk> reConfigChunk{ RTC::SCTP::ReConfigChunk::Factory(
			buffer.data(), buffer.size()) };

		test.HandleReceivedReConfigChunk(reConfigChunk.get());

		REQUIRE(test.associationListener.HasErrored() == true);
		REQUIRE(test.associationListener.HasSentPackets() == false);
	}

	SECTION("a chunk with invalid parameters returns an error")
	{
		TestStreamResetHandler test;

		std::vector<uint8_t> buffer(test.sctpOptions.mtu);

		const std::unique_ptr<RTC::SCTP::ReConfigChunk> reConfigChunk{ RTC::SCTP::ReConfigChunk::Factory(
			buffer.data(), buffer.size()) };

		// Two RTC::SCTP::OutgoingSsnResetRequestParameter in a RE-CONFIG is not valid.
		for (const uint32_t reqSeqNbr : { 1u, 2u })
		{
			auto* parameter =
			  reConfigChunk->BuildParameterInPlace<RTC::SCTP::OutgoingSsnResetRequestParameter>();

			parameter->SetReconfigurationRequestSequenceNumber(reqSeqNbr);
			parameter->SetReconfigurationResponseSequenceNumber(10);
			parameter->SetSenderLastAssignedTsn(RemoteInitialTsn);
			parameter->AddStreamId(static_cast<uint16_t>(reqSeqNbr));
			parameter->Consolidate();
		}

		test.HandleReceivedReConfigChunk(reConfigChunk.get());

		REQUIRE(test.associationListener.HasErrored() == true);
		REQUIRE(test.associationListener.HasSentPackets() == false);
	}

	SECTION("sends an outgoing request directly")
	{
		TestStreamResetHandler test;

		const std::vector<uint16_t> streamIds{ 42 };

		test.streamResetHandler.ResetStreams(streamIds);

		REQUIRE(test.streamResetHandler.ShouldSendStreamResetRequest() == true);

		const auto packet = test.tcbContext.CreatePacket();

		test.streamResetHandler.AddStreamResetRequest(packet.get());

		const auto* reConfigChunk = packet->GetFirstChunkOfType<RTC::SCTP::ReConfigChunk>();

		REQUIRE(reConfigChunk);

		const auto* parameter =
		  reConfigChunk->GetFirstParameterOfType<RTC::SCTP::OutgoingSsnResetRequestParameter>();

		REQUIRE(parameter);
		REQUIRE(parameter->GetStreamIds() == std::vector<uint16_t>{ 42 });
	}

	SECTION("resets multiple streams in one request")
	{
		TestStreamResetHandler test;

		test.streamResetHandler.ResetStreams(std::vector<uint16_t>{ 42 });
		test.streamResetHandler.ResetStreams(std::vector<uint16_t>{ 43, 44, 41 });
		test.streamResetHandler.ResetStreams(std::vector<uint16_t>{ 42, 40 });

		REQUIRE(test.streamResetHandler.ShouldSendStreamResetRequest() == true);

		const auto packet = test.tcbContext.CreatePacket();

		test.streamResetHandler.AddStreamResetRequest(packet.get());

		const auto* reConfigChunk = packet->GetFirstChunkOfType<RTC::SCTP::ReConfigChunk>();

		REQUIRE(reConfigChunk);

		const auto* parameter =
		  reConfigChunk->GetFirstParameterOfType<RTC::SCTP::OutgoingSsnResetRequestParameter>();

		REQUIRE(parameter);

		auto streamIds = parameter->GetStreamIds();

		std::ranges::sort(streamIds);

		REQUIRE(streamIds == std::vector<uint16_t>{ 40, 41, 42, 43, 44 });
	}

	SECTION("commits the reset on a positive response")
	{
		TestStreamResetHandler test;

		test.streamResetHandler.ResetStreams(std::vector<uint16_t>{ 42 });

		REQUIRE(test.streamResetHandler.ShouldSendStreamResetRequest() == true);

		const auto packet = test.tcbContext.CreatePacket();

		test.streamResetHandler.AddStreamResetRequest(packet.get());

		const auto* parameter =
		  packet->GetFirstChunkOfType<RTC::SCTP::ReConfigChunk>()
		    ->GetFirstParameterOfType<RTC::SCTP::OutgoingSsnResetRequestParameter>();

		REQUIRE(parameter);

		const uint32_t reqSeqNbr = parameter->GetReconfigurationRequestSequenceNumber();

		// Build a RE-CONFIG response with a successful result.
		std::vector<uint8_t> responseBuffer(test.sctpOptions.mtu);
		const std::unique_ptr<RTC::SCTP::ReConfigChunk> responseChunk{ RTC::SCTP::ReConfigChunk::Factory(
			responseBuffer.data(), responseBuffer.size()) };

		auto* response =
		  responseChunk->BuildParameterInPlace<RTC::SCTP::ReconfigurationResponseParameter>();

		response->SetReconfigurationResponseSequenceNumber(reqSeqNbr);
		response->SetResult(RTC::SCTP::ReconfigurationResponseParameter::Result::SUCCESS_PERFORMED);
		response->Consolidate();

		test.HandleReceivedReConfigChunk(responseChunk.get());

		REQUIRE(test.associationListener.HasStreamsResetPerformedForStreamId(42) == true);
	}

	SECTION("rolls back the reset on an error response")
	{
		TestStreamResetHandler test;

		test.streamResetHandler.ResetStreams(std::vector<uint16_t>{ 42 });

		REQUIRE(test.streamResetHandler.ShouldSendStreamResetRequest() == true);

		const auto packet = test.tcbContext.CreatePacket();

		test.streamResetHandler.AddStreamResetRequest(packet.get());

		const auto* parameter =
		  packet->GetFirstChunkOfType<RTC::SCTP::ReConfigChunk>()
		    ->GetFirstParameterOfType<RTC::SCTP::OutgoingSsnResetRequestParameter>();

		REQUIRE(parameter);

		const uint32_t reqSeqNbr = parameter->GetReconfigurationRequestSequenceNumber();

		// Build a RE-CONFIG response with an error result.
		std::vector<uint8_t> responseBuffer(test.sctpOptions.mtu);
		const std::unique_ptr<RTC::SCTP::ReConfigChunk> responseChunk{ RTC::SCTP::ReConfigChunk::Factory(
			responseBuffer.data(), responseBuffer.size()) };

		auto* response =
		  responseChunk->BuildParameterInPlace<RTC::SCTP::ReconfigurationResponseParameter>();

		response->SetReconfigurationResponseSequenceNumber(reqSeqNbr);
		response->SetResult(
		  RTC::SCTP::ReconfigurationResponseParameter::Result::ERROR_BAD_SEQUENCE_NUMBER);
		response->Consolidate();

		test.HandleReceivedReConfigChunk(responseChunk.get());

		REQUIRE(test.associationListener.HasStreamsResetFailedForStreamId(42) == true);
	}

	SECTION("an incoming stream reset request returns nothing performed")
	{
		TestStreamResetHandler test;

		std::vector<uint8_t> buffer(test.sctpOptions.mtu);
		const std::unique_ptr<RTC::SCTP::ReConfigChunk> reConfigChunk{ RTC::SCTP::ReConfigChunk::Factory(
			buffer.data(), buffer.size()) };

		auto* parameter =
		  reConfigChunk->BuildParameterInPlace<RTC::SCTP::IncomingSsnResetRequestParameter>();

		parameter->SetReconfigurationRequestSequenceNumber(RemoteInitialTsn);
		parameter->AddStreamId(1);
		parameter->Consolidate();

		test.HandleReceivedReConfigChunk(reConfigChunk.get());

		// A RE-CONFIG response is sent back.
		const auto sentBuffer = test.associationListener.ConsumeFirstSentPacket();

		REQUIRE(!sentBuffer.empty());

		const std::unique_ptr<RTC::SCTP::Packet> sentPacket{ RTC::SCTP::Packet::Parse(
			sentBuffer.data(), sentBuffer.size()) };

		REQUIRE(sentPacket);

		const auto* response = sentPacket->GetFirstChunkOfType<RTC::SCTP::ReConfigChunk>()
		                         ->GetFirstParameterOfType<RTC::SCTP::ReconfigurationResponseParameter>();

		REQUIRE(response);
		REQUIRE(
		  response->GetResult() ==
		  RTC::SCTP::ReconfigurationResponseParameter::Result::SUCCESS_NOTHING_TO_DO);
	}

	// @see https://github.com/versatica/mediasoup/security/advisories/GHSA-rq7g-r9qr-rwpq
	SECTION("received Forward-TSN chunks are bounded during deferred reset processing")
	{
		TestStreamResetHandler test;

		// Makes the peer request an outgoing stream reset with the given sender's
		// last assigned TSN.
		auto handleReceivedOutgoingSsnResetRequest =
		  [&test](uint32_t reqSeqNbr, uint32_t senderLastAssignedTsn) -> void
		{
			std::vector<uint8_t> buffer(test.sctpOptions.mtu);

			const std::unique_ptr<RTC::SCTP::ReConfigChunk> reConfigChunk{
				RTC::SCTP::ReConfigChunk::Factory(buffer.data(), buffer.size())
			};

			auto* parameter =
			  reConfigChunk->BuildParameterInPlace<RTC::SCTP::OutgoingSsnResetRequestParameter>();

			parameter->SetReconfigurationRequestSequenceNumber(reqSeqNbr);
			parameter->SetReconfigurationResponseSequenceNumber(0);
			parameter->SetSenderLastAssignedTsn(senderLastAssignedTsn);
			parameter->AddStreamId(1);
			parameter->Consolidate();

			test.HandleReceivedReConfigChunk(reConfigChunk.get());
		};

		// Returns the result of the RE-CONFIG response sent back to the peer.
		auto consumeSentReConfigResponseResult =
		  [&test]() -> RTC::SCTP::ReconfigurationResponseParameter::Result
		{
			const auto sentBuffer = test.associationListener.ConsumeFirstSentPacket();

			REQUIRE(!sentBuffer.empty());

			const std::unique_ptr<RTC::SCTP::Packet> sentPacket{ RTC::SCTP::Packet::Parse(
				sentBuffer.data(), sentBuffer.size()) };

			REQUIRE(sentPacket);

			const auto* response =
			  sentPacket->GetFirstChunkOfType<RTC::SCTP::ReConfigChunk>()
			    ->GetFirstParameterOfType<RTC::SCTP::ReconfigurationResponseParameter>();

			REQUIRE(response);

			return response->GetResult();
		};

		// Feeds a received packet carrying a single FORWARD-TSN chunk to the receive
		// side, doing exactly what RTC::SCTP::Association::HandleReceivedAnyForwardTsnChunk()
		// and then RTC::SCTP::Association::ReceiveSctpData() do.
		auto handleReceivedForwardTsn =
		  [&test](
		    uint32_t newCumulativeTsn,
		    std::span<const RTC::SCTP::AnyForwardTsnChunk::SkippedStream> skippedStreams) -> void
		{
			const RTC::SCTP::AssociationListenerDeferrer::ScopedDeferrer deferrer(
			  test.associationListenerDeferrer);

			if (test.dataTracker.HandleForwardTsn(newCumulativeTsn))
			{
				test.reassemblyQueue.HandleForwardTsn(newCumulativeTsn, skippedStreams);
			}

			test.streamResetHandler.MayLeaveDeferredReset();
		};

		// The number of skipped streams that fit in a single MTU sized FORWARD-TSN
		// chunk: the SCTP common header (12 bytes), the chunk header (4 bytes) and
		// the new cumulative TSN (4 bytes), and then 4 bytes per skipped stream.
		const size_t skippedStreamsPerChunk = (test.sctpOptions.mtu - 12 - 4 - 4) / 4;

		std::vector<RTC::SCTP::AnyForwardTsnChunk::SkippedStream> skippedStreams;

		skippedStreams.reserve(skippedStreamsPerChunk);

		for (size_t i{ 0 }; i < skippedStreamsPerChunk; ++i)
		{
			skippedStreams.emplace_back(/*streamId*/ 1, /*ssn*/ static_cast<uint16_t>(i));
		}

		// The peer requests a stream reset with a sender's last assigned TSN that
		// has not been reached yet, so the receive side enters deferred reset
		// processing and replies "in progress".
		handleReceivedOutgoingSsnResetRequest(
		  /*reqSeqNbr*/ RemoteInitialTsn, /*senderLastAssignedTsn*/ RemoteInitialTsn + 1);

		REQUIRE(
		  consumeSentReConfigResponseResult() ==
		  RTC::SCTP::ReconfigurationResponseParameter::Result::IN_PROGRESS);

		// The peer now sends FORWARD-TSN chunks with ever increasing new cumulative
		// TSNs and never sends the follow-up reset request. Every such chunk is
		// beyond the sender's last assigned TSN, so it takes the deferred branch of
		// RTC::SCTP::ReassemblyQueue::HandleForwardTsn().
		//
		// The cumulative ack TSN reaches the sender's last assigned TSN with the
		// very first of them, so deferred reset processing must end there and the
		// queued bytes must never grow beyond the receive buffer.
		for (uint32_t tsn{ RemoteInitialTsn + 2 }; tsn < RemoteInitialTsn + 20000; ++tsn)
		{
			handleReceivedForwardTsn(tsn, skippedStreams);
		}

		REQUIRE(test.reassemblyQueue.GetQueuedBytes() <= test.sctpOptions.maxReceiverWindowBufferSize);
	}
}
