#define MS_CLASS "RTC::SCTP::Packet"
// #define MS_LOG_DEV_LEVEL 3

#include "RTC/SCTP/association/StateCookie.hpp"
#include "Logger.hpp"
#include "MediaSoupErrors.hpp"
#include <openssl/crypto.h>
#include <cstring> // std::memcpy()
#include <string_view>

namespace RTC
{
	namespace SCTP
	{
		/* Class methods. */

		StateCookie* StateCookie::Parse(const uint8_t* buffer, size_t bufferLength)
		{
			MS_TRACE();

			if (!StateCookie::IsMediasoupStateCookie(buffer, bufferLength))
			{
				MS_WARN_TAG(sctp, "not a StateCookie generated by mediasoup");

				return nullptr;
			}

			// Reject cookies with a zero verification tag. A zero tag is never
			// generated by us, so this limits the effect of a tampered/forged cookie.
			if (Utils::Byte::Get4Bytes(buffer, 8) == 0 || Utils::Byte::Get4Bytes(buffer, 12) == 0)
			{
				MS_WARN_TAG(sctp, "invalid StateCookie, verification tag is zero");

				return nullptr;
			}

			auto* stateCookie = new StateCookie(const_cast<uint8_t*>(buffer), bufferLength);

			return stateCookie;
		}

		StateCookie* StateCookie::Factory(
		  uint8_t* buffer,
		  size_t bufferLength,
		  uint32_t localVerificationTag,
		  uint32_t remoteVerificationTag,
		  uint32_t localInitialTsn,
		  uint32_t remoteInitialTsn,
		  uint32_t remoteAdvertisedReceiverWindowCredit,
		  uint64_t tieTag,
		  const Capabilities& remoteCapabilities,
		  int64_t creationTimestampUs,
		  const uint8_t* macKey,
		  size_t macKeyLength)
		{
			MS_TRACE();

			// This may throw.
			StateCookie::Write(
			  buffer,
			  bufferLength,
			  localVerificationTag,
			  remoteVerificationTag,
			  localInitialTsn,
			  remoteInitialTsn,
			  remoteAdvertisedReceiverWindowCredit,
			  tieTag,
			  remoteCapabilities,
			  creationTimestampUs,
			  macKey,
			  macKeyLength);

			const size_t cookieLength = macKey != nullptr ? StateCookie::AuthenticatedStateCookieLength
			                                              : StateCookie::StateCookieLength;

			return new StateCookie(buffer, cookieLength);
		}

		void StateCookie::Write(
		  uint8_t* buffer,
		  size_t bufferLength,
		  uint32_t localVerificationTag,
		  uint32_t remoteVerificationTag,
		  uint32_t localInitialTsn,
		  uint32_t remoteInitialTsn,
		  uint32_t remoteAdvertisedReceiverWindowCredit,
		  uint64_t tieTag,
		  const Capabilities& remoteCapabilities,
		  int64_t creationTimestampUs,
		  const uint8_t* macKey,
		  size_t macKeyLength)
		{
			MS_TRACE();

			const bool authenticate = macKey != nullptr;
			const size_t cookieLength =
			  authenticate ? StateCookie::AuthenticatedStateCookieLength : StateCookie::StateCookieLength;

			if (bufferLength < cookieLength)
			{
				MS_THROW_TYPE_ERROR("buffer too small");
			}

			Utils::Byte::Set8Bytes(buffer, 0, StateCookie::Magic1);
			Utils::Byte::Set4Bytes(buffer, 8, localVerificationTag);
			Utils::Byte::Set4Bytes(buffer, 12, remoteVerificationTag);
			Utils::Byte::Set4Bytes(buffer, 16, localInitialTsn);
			Utils::Byte::Set4Bytes(buffer, 20, remoteInitialTsn);
			Utils::Byte::Set4Bytes(buffer, 24, remoteAdvertisedReceiverWindowCredit);
			Utils::Byte::Set8Bytes(buffer, 28, tieTag);

			auto* remoteCapabilitiesField =
			  reinterpret_cast<RemoteCapabilitiesField*>(buffer + StateCookie::RemoteCapabilitiesOffset);

			remoteCapabilitiesField->reserved   = 0;
			remoteCapabilitiesField->bitA       = remoteCapabilities.partialReliability;
			remoteCapabilitiesField->bitB       = remoteCapabilities.messageInterleaving;
			remoteCapabilitiesField->bitC       = remoteCapabilities.reConfig;
			remoteCapabilitiesField->unusedBits = 0;
			remoteCapabilitiesField->magic2     = htons(StateCookie::Magic2);
			remoteCapabilitiesField->zeroChecksumAlternateErrorDetectionMethod =
			  htonl(static_cast<uint32_t>(remoteCapabilities.zeroChecksumAlternateErrorDetectionMethod));
			remoteCapabilitiesField->maxOutboundStreams = htons(remoteCapabilities.maxOutboundStreams);
			remoteCapabilitiesField->maxInboundStreams  = htons(remoteCapabilities.maxInboundStreams);

			if (!authenticate)
			{
				return;
			}

			// Append the creation timestamp and the MAC computed over all preceding
			// bytes (including the timestamp).
			//
			// @see RFC 9260 section 5.1.3.
			Utils::Byte::Set8Bytes(
			  buffer, StateCookie::TimestampOffset, static_cast<uint64_t>(creationTimestampUs));

			const uint8_t* mac = Utils::Crypto::GetHmacSha1(
			  reinterpret_cast<const char*>(macKey), macKeyLength, buffer, StateCookie::MacOffset);

			std::memcpy(buffer + StateCookie::MacOffset, mac, StateCookie::MacLength);
		}

		bool StateCookie::IsMediasoupStateCookie(const uint8_t* buffer, size_t bufferLength)
		{
			MS_TRACE();

			if (bufferLength != StateCookie::StateCookieLength && bufferLength != StateCookie::AuthenticatedStateCookieLength)
			{
				return false;
			}

			if (Utils::Byte::Get8Bytes(buffer, 0) != StateCookie::Magic1)
			{
				return false;
			}

			const auto* remoteCapabilitiesField = reinterpret_cast<const RemoteCapabilitiesField*>(
			  buffer + StateCookie::RemoteCapabilitiesOffset);

			if (ntohs(remoteCapabilitiesField->magic2) != StateCookie::Magic2)
			{
				return false;
			}

			return true;
		}

		bool StateCookie::VerifyMac(
		  const uint8_t* buffer, size_t bufferLength, const uint8_t* macKey, size_t macKeyLength)
		{
			MS_TRACE();

			// An authenticated cookie has a fixed length.
			if (bufferLength != StateCookie::AuthenticatedStateCookieLength)
			{
				return false;
			}

			// Recompute the MAC over all bytes preceding the MAC field.
			const uint8_t* expectedMac = Utils::Crypto::GetHmacSha1(
			  reinterpret_cast<const char*>(macKey), macKeyLength, buffer, StateCookie::MacOffset);

			// NOTE: Use `CRYPTO_memcmp()` to have constant time memory comparison.
			// See https://github.com/versatica/mediasoup/security/advisories/GHSA-xvjj-6cm4-ppgq
			return CRYPTO_memcmp(buffer + StateCookie::MacOffset, expectedMac, StateCookie::MacLength) == 0;
		}

		Types::SctpImplementation StateCookie::DetermineSctpImplementation(
		  const uint8_t* buffer, size_t bufferLength)
		{
			MS_TRACE();

			if (bufferLength < StateCookie::Magic1Length)
			{
				return Types::SctpImplementation::UNKNOWN;
			}

			const std::string_view magic1(reinterpret_cast<const char*>(buffer), StateCookie::Magic1Length);

			if (magic1 == "msworker")
			{
				return Types::SctpImplementation::MEDIASOUP;
			}
			else if (magic1 == "dcSCTP00")
			{
				return Types::SctpImplementation::DCSCTP;
			}
			else if (magic1 == "KAME-BSD")
			{
				return Types::SctpImplementation::USRSCTP;
			}
			else
			{
				return Types::SctpImplementation::UNKNOWN;
			}
		}

		/* Instance methods. */

		StateCookie::StateCookie(uint8_t* buffer, size_t bufferLength)
		  : Serializable(buffer, bufferLength)
		{
			MS_TRACE();

			// The cookie length is determined by whether it carries a MAC. When the
			// cookie is cloned into a larger buffer, `CloneInto()` will set the
			// proper length afterwards.
			SetLength(
			  bufferLength == StateCookie::AuthenticatedStateCookieLength
			    ? StateCookie::AuthenticatedStateCookieLength
					: StateCookie::StateCookieLength);
		}

		StateCookie::~StateCookie()
		{
			MS_TRACE();
		}

		void StateCookie::Dump(int indentation) const
		{
			MS_TRACE();

			const auto remoteCapabilities = GetRemoteCapabilities();

			MS_DUMP_CLEAN(indentation, "<SCTP::StateCookie>");

			MS_DUMP_CLEAN(indentation, "  length: %zu (buffer length: %zu)", GetLength(), GetBufferLength());
			MS_DUMP_CLEAN(indentation, "  local verification tag: %" PRIu32, GetLocalVerificationTag());
			MS_DUMP_CLEAN(indentation, "  remote verification tag: %" PRIu32, GetRemoteVerificationTag());
			MS_DUMP_CLEAN(indentation, "  local initial tsn: %" PRIu32, GetLocalInitialTsn());
			MS_DUMP_CLEAN(indentation, "  remote initial tsn: %" PRIu32, GetRemoteInitialTsn());
			MS_DUMP_CLEAN(
			  indentation,
			  "  remote advertised receiver window credit: %" PRIu32,
			  GetRemoteAdvertisedReceiverWindowCredit());
			MS_DUMP_CLEAN(indentation, "  tie-tag: %" PRIu64, GetTieTag());
			MS_DUMP_CLEAN(indentation, "  authenticated: %s", IsAuthenticated() ? "yes" : "no");

			if (IsAuthenticated())
			{
				MS_DUMP_CLEAN(indentation, "  creation timestamp (us): %" PRIi64, GetCreationTimestampUs());
			}

			MS_DUMP_CLEAN(indentation, "  remote capabilities:");
			remoteCapabilities.Dump(indentation + 1);

			MS_DUMP_CLEAN(indentation, "</SCTP::StateCookie>");
		}

		StateCookie* StateCookie::Clone(uint8_t* buffer, size_t bufferLength) const
		{
			MS_TRACE();

			auto* clonedStateCookie = new StateCookie(buffer, bufferLength);

			Serializable::CloneInto(clonedStateCookie);

			return clonedStateCookie;
		}

		Capabilities StateCookie::GetRemoteCapabilities() const
		{
			MS_TRACE();

			const auto* remoteCapabilitiesField = GetRemoteCapabilitiesField();

			Capabilities remoteCapabilities;

			remoteCapabilities.maxOutboundStreams  = ntohs(remoteCapabilitiesField->maxOutboundStreams);
			remoteCapabilities.maxInboundStreams   = ntohs(remoteCapabilitiesField->maxInboundStreams);
			remoteCapabilities.partialReliability  = remoteCapabilitiesField->bitA;
			remoteCapabilities.messageInterleaving = remoteCapabilitiesField->bitB;
			remoteCapabilities.reConfig            = remoteCapabilitiesField->bitC;

			// Only keep the zero checksum method if it's a value we know about.
			// Anything else (unknown/future method or a tampered cookie) is treated
			// as none.
			const uint32_t zeroChecksumAlternateErrorDetectionMethod =
			  ntohl(remoteCapabilitiesField->zeroChecksumAlternateErrorDetectionMethod);

			remoteCapabilities.zeroChecksumAlternateErrorDetectionMethod =
			  zeroChecksumAlternateErrorDetectionMethod ==
			      static_cast<uint32_t>(
			        ZeroChecksumAcceptableParameter::AlternateErrorDetectionMethod::SCTP_OVER_DTLS)
			    ? ZeroChecksumAcceptableParameter::AlternateErrorDetectionMethod::SCTP_OVER_DTLS
			    : ZeroChecksumAcceptableParameter::AlternateErrorDetectionMethod::NONE;

			// NOTE: No need to std::move(). Copy elision (RVO) is used for free in GCC
			// and clang in C++17 or higher.
			return remoteCapabilities;
		}
	} // namespace SCTP
} // namespace RTC
