--
-- == VECTOR LIBRARY ==
-- == by TehRealSalt ==
--
-- An SRB2 Lua implementation for Vector2 and Vector3 data types.
-- Physics have never been easier!
--

local OURVERSION = 3;
local REPLACEVERSION = rawget(_G, "VectorLibVersion");

if (REPLACEVERSION != nil and OURVERSION <= REPLACEVERSION)
	-- Already loaded.
	return;
end

rawset(_G, "VectorLibVersion", OURVERSION);

-- These help with determining signs.
local function sign(x)
	if (x == 0)
		return 0;
	elseif (x < 0)
		return -1
	else
		return 1;
	end
end

local function FixedSign(x)
	return Sign(x) * FRACUNIT;
end

local function clamp(x, l, h)
	return min(max(x, l), h);
end

rawset(_G, "sign", sign);
rawset(_G, "FixedSign", FixedSign);
rawset(_G, "clamp", clamp);

local Vector2 = {};
local Vector3 = {};

Vector2.Meta = {};
Vector3.Meta = {};

--
-- VECTOR2
--

Vector2.New = function(vecX, vecY)
	local vector = {};
	setmetatable(vector, Vector2.Meta);

	vector.x = vecX;
	vector.y = vecY;

	return vector;
end

Vector2.Copy = function(vector)
	return Vector2.New(vector.x, vector.y);
end

Vector2.Zero = Vector2.New(0, 0);

Vector2.Convert = function(x)
	if (getmetatable(x) == Vector2.Meta)
		-- It is a Vector3.
		return x;
	end

	local xType = type(x);
	if (xType == "number")
		-- Convert fixed into vector, for operations.
		return Vector2.New(x, x);
	elseif (xType == "table")
		-- Generic table, turn its values into a Vector3.
		local newVecX = 0;
		local newVecY = 0;

		if (x.x or x.y)
			-- Has proper named fields.
			if (type(x.x) == "number")
				newVecX = x.x;
			end

			if (type(x.y) == "number")
				newVecY = x.y;
			end
		else
			-- Unnamed / weird named fields,
			-- just use them sequentially.
			local i = 0;
			for _,v in ipairs(x) do
				if (i == 0)
					newVecX = v;
				elseif (i == 1)
					newVecY = v;
				else
					break;
				end

				i = $1 + 1;
			end
		end

		return Vector2.New(newVecX, newVecY);
	end

	error("Vector2.Convert received unexpected type '" + xType + "'.");
	return nil;
end

Vector2.Add = function(a, b)
	local vectorA = Vector2.Convert(a);
	local vectorB = Vector2.Convert(b);

	return Vector2.New(
		vectorA.x + vectorB.x,
		vectorA.y + vectorB.y
	);
end

Vector2.Sub = function(a, b)
	local vectorA = Vector2.Convert(a);
	local vectorB = Vector2.Convert(b);

	return Vector2.New(
		vectorA.x - vectorB.x,
		vectorA.y - vectorB.y
	);
end

Vector2.Mul = function(a, b)
	local vectorA = Vector2.Convert(a);
	local vectorB = Vector2.Convert(b);

	return Vector2.New(
		FixedMul(vectorA.x, vectorB.x),
		FixedMul(vectorA.y, vectorB.y)
	);
end

Vector2.Div = function(a, b)
	local vectorA = Vector2.Convert(a);
	local vectorB = Vector2.Convert(b);

	return Vector2.New(
		FixedDiv(vectorA.x, vectorB.x),
		FixedDiv(vectorA.y, vectorB.y)
	);
end

Vector2.IntegerMul = function(a, b)
	local vectorA = Vector2.Convert(a);
	local vectorB = Vector2.Convert(b);

	return Vector2.New(
		vectorA.x * vectorB.x,
		vectorA.y * vectorB.y
	);
end

Vector2.IntegerDiv = function(a, b)
	local vectorA = Vector2.Convert(a);
	local vectorB = Vector2.Convert(b);

	return Vector2.New(
		vectorA.x / vectorB.x,
		vectorA.y / vectorB.y
	);
end

Vector2.Negate = function(a)
	local vectorA = Vector2.Convert(a);
	local vectorB = Vector2.Convert(b);

	return Vector2.New(
		-vectorA.x,
		-vectorA.y
	);
end

Vector2.Equal = function(a, b)
	local vectorA = Vector2.Convert(a);
	local vectorB = Vector2.Convert(b);

	if (vectorA.x == vectorB.x)
	and (vectorA.y == vectorB.y)
		return true;
	end

	return false;
end

Vector2.Invert = function(vector)
	return Vector2.Div(FRACUNIT, vector);
end

Vector2.Abs = function(vector)
	return Vector2.New(
		abs(vector.x),
		abs(vector.y)
	);
end

Vector2.Sign = function(vector)
	return Vector2.New(
		FixedSign(vector.x),
		FixedSign(vector.y)
	);
end

Vector2.IntegerSign = function(vector)
	return Vector2.New(
		sign(vector.x),
		sign(vector.y)
	);
end

Vector2.Floor = function(vector)
	return Vector2.New(
		FixedFloor(vector.x),
		FixedFloor(vector.y)
	);
end

Vector2.Ceil = function(vector)
	return Vector2.New(
		FixedCeil(vector.x),
		FixedCeil(vector.y)
	);
end

Vector2.Round = function(vector)
	return Vector2.New(
		FixedRound(vector.x),
		FixedRound(vector.y)
	);
end

Vector2.Trunc = function(vector)
	return Vector2.New(
		FixedTrunc(vector.x),
		FixedTrunc(vector.y)
	);
end

Vector2.Length = function(vector)
	return R_PointToDist2(0, 0, vector.x, vector.y);
end

Vector2.Angle = function(vector)
	return R_PointToAngle2(0, 0, vector.x, vector.y);
end

Vector2.DistanceTo = function(vectorA, vectorB)
	return R_PointToDist2(vectorA.x, vectorA.y, vectorB.x, vectorB.y);
end

Vector2.AngleTo = function(vectorA, vectorB)
	return R_PointToAngle2(vectorA.x, vectorA.y, vectorB.x, vectorB.y);
end

Vector2.Normalize = function(vector)
	local length = Vector2.Length(vector);

	if (length > 0)
		return Vector2.Div(vector, length);
	end

	return Vector2.New(0, 0);
end

Vector2.DirectionTo = function(vectorA, vectorB)
	return Vector2.Normalize(Vector2.Sub(vectorB, vectorA));
end

Vector2.Clamp = function(vector, maxLength)
	local length = Vector2.Length(vector);

	if (length > maxLength)
		return Vector2.Mul(Vector2.Normalized(vector), maxLength);
	end

	return Vector2.Copy(vector);
end

Vector2.Rotate = function(vector, angle)
	local sine = sin(angle);
	local cosine = cos(angle);

	return Vector2.New(
		FixedMul(vector.x, cosine) - FixedMul(vector.y, sine),
		FixedMul(vector.x, sine) + FixedMul(vector.y, cosine)
	);
end

Vector2.Dot = function(vectorA, vectorB)
	return FixedMul(vectorA.x, vectorB.x) + FixedMul(vectorA.y, vectorB.y);
end

Vector2.Cross = function(vectorA, vectorB)
	return FixedMul(vectorA.x, vectorB.y) - FixedMul(vectorA.y, vectorB.x);
end

Vector2.MobjPositionVector = function(mobj)
	return Vector2.New(mobj.x, mobj.y);
end

Vector2.MobjMomentumVector = function(mobj)
	return Vector2.New(mobj.momx, mobj.momy);
end

Vector2.SetMobjMomentum = function(mobj, vector)
	if not (mobj and mobj.valid)
		return;
	end

	mobj.momx = vector.x;
	mobj.momy = vector.y;
end

Vector2.Meta = {
	__add = Vector2.Add,
	__sub = Vector2.Sub,
	__mul = Vector2.IntegerMul,
	__div = Vector2.IntegerDiv,
	__unm = Vector2.Negate,
	__eq = Vector2.Equal,
};
registerMetatable(Vector2.Meta);

rawset(_G, "Vector2", Vector2);

--
-- VECTOR3
--

Vector3.New = function(vecX, vecY, vecZ)
	local vector = {};
	setmetatable(vector, Vector3.Meta);

	vector.x = vecX;
	vector.y = vecY;
	vector.z = vecZ;

	return vector;
end

Vector3.Copy = function(vector)
	return Vector3.New(
		vector.x,
		vector.y,
		vector.z
	);
end

Vector3.Zero = Vector3.New(0, 0, 0);

Vector3.ToVector2 = function(vector)
	return Vector2.New(vector.x, vector.y);
end

Vector3.FromVector2 = function(vector)
	return Vector3.New(vector.x, vector.y, 0);
end

Vector3.Convert = function(x)
	if (getmetatable(x) == Vector3.Meta)
		-- It is a Vector3.
		return x;
	end

	local xType = type(x);
	if (xType == "number")
		-- Convert fixed into vector, for operations.
		return Vector3.New(x, x, x);
	elseif (xType == "table")
		-- Generic table, turn its values into a Vector3.
		local newVecX = 0;
		local newVecY = 0;
		local newVecZ = 0;

		if (x.x or x.y or x.z)
			-- Has proper named fields.
			if (type(x.x) == "number")
				newVecX = x.x;
			end

			if (type(x.y) == "number")
				newVecY = x.y;
			end

			if (type(x.z) == "number")
				newVecZ = x.z;
			end
		else
			-- Unnamed / weird named fields,
			-- just use them sequentially.
			local i = 0;
			for _,v in ipairs(x) do
				if (i == 0)
					newVecX = v;
				elseif (i == 1)
					newVecY = v;
				elseif (i == 2)
					newVecZ = v;
				else
					break;
				end

				i = $1 + 1;
			end
		end

		return Vector3.New(newVecX, newVecY, newVecZ);
	end

	error("Vector3.Convert received unexpected type '" + xType + "'.");
	return nil;
end

Vector3.Add = function(a, b)
	local vectorA = Vector3.Convert(a);
	local vectorB = Vector3.Convert(b);

	return Vector3.New(
		vectorA.x + vectorB.x,
		vectorA.y + vectorB.y,
		vectorA.z + vectorB.z
	);
end

Vector3.Sub = function(a, b)
	local vectorA = Vector3.Convert(a);
	local vectorB = Vector3.Convert(b);

	return Vector3.New(
		vectorA.x - vectorB.x,
		vectorA.y - vectorB.y,
		vectorA.z - vectorB.z
	);
end

Vector3.Mul = function(a, b)
	local vectorA = Vector3.Convert(a);
	local vectorB = Vector3.Convert(b);

	return Vector3.New(
		FixedMul(vectorA.x, vectorB.x),
		FixedMul(vectorA.y, vectorB.y),
		FixedMul(vectorA.z, vectorB.z)
	);
end

Vector3.Div = function(a, b)
	local vectorA = Vector3.Convert(a);
	local vectorB = Vector3.Convert(b);

	return Vector3.New(
		FixedDiv(vectorA.x, vectorB.x),
		FixedDiv(vectorA.y, vectorB.y),
		FixedDiv(vectorA.z, vectorB.z)
	);
end

Vector3.IntegerMul = function(a, b)
	local vectorA = Vector3.Convert(a);
	local vectorB = Vector3.Convert(b);

	return Vector3.New(
		vectorA.x * vectorB.x,
		vectorA.y * vectorB.y,
		vectorA.z * vectorB.z
	);
end

Vector3.IntegerDiv = function(a, b)
	local vectorA = Vector3.Convert(a);
	local vectorB = Vector3.Convert(b);

	return Vector3.New(
		vectorA.x / vectorB.x,
		vectorA.y / vectorB.y,
		vectorA.z / vectorB.z
	);
end

Vector3.Negate = function(a)
	local vectorA = Vector3.Convert(a);

	return Vector3.New(
		-vectorA.x,
		-vectorA.y,
		-vectorA.z
	);
end

Vector3.Equal = function(a, b)
	local vectorA = Vector3.Convert(a);
	local vectorB = Vector3.Convert(b);

	if (vectorA.x == vectorB.x)
	and (vectorA.y == vectorB.y)
	and (vectorA.z == vectorB.z)
		return true;
	end

	return false;
end

Vector3.Invert = function(vector)
	return Vector3.Div(FRACUNIT, vector);
end

Vector3.Abs = function(vector)
	return Vector3.New(
		abs(vector.x),
		abs(vector.y),
		abs(vector.z)
	);
end

Vector3.Sign = function(vector)
	return Vector3.New(
		FixedSign(vector.x),
		FixedSign(vector.y),
		FixedSign(vector.z)
	);
end

Vector3.IntegerSign = function(vector)
	return Vector3.New(
		sign(vector.x),
		sign(vector.y),
		sign(vector.z)
	);
end

Vector3.Floor = function(vector)
	return Vector3.New(
		FixedFloor(vector.x),
		FixedFloor(vector.y),
		FixedFloor(vector.z)
	);
end

Vector3.Ceil = function(vector)
	return Vector3.New(
		FixedCeil(vector.x),
		FixedCeil(vector.y),
		FixedCeil(vector.z)
	);
end

Vector3.Round = function(vector)
	return Vector3.New(
		FixedRound(vector.x),
		FixedRound(vector.y),
		FixedRound(vector.z)
	);
end

Vector3.Trunc = function(vector)
	return Vector3.New(
		FixedTrunc(vector.x),
		FixedTrunc(vector.y),
		FixedTrunc(vector.z)
	);
end

Vector3.Length = function(vector)
	local xyLength = R_PointToDist2(0, 0, vector.x, vector.y);
	return R_PointToDist2(0, 0, xyLength, vector.z);
end

Vector3.Aiming = function(vector)
	local xyLength = R_PointToDist2(0, 0, vector.x, vector.y);
	return R_PointToAngle2(0, 0, xyLength, vector.z);
end

Vector3.Angles = function(vector)
	return {
		angle = Vector2.Angle(vector),
		aiming = Vector3.Aiming(vector),
	};
end

Vector3.DistanceTo = function(vectorA, vectorB)
	return Vector3.Length(Vector3.Sub(vectorA, vectorB));
end

Vector3.AngleTo = function(vectorA, vectorB)
	return Vector3.Angles(Vector3.Sub(vectorA, vectorB));
end

Vector3.Normalize = function(vector)
	local length = Vector3.Length(vector);

	if (length > 0)
		return Vector3.Div(vector, length);
	end

	return Vector3.New(0, 0, 0);
end

Vector3.DirectionTo = function(vectorA, vectorB)
	return Vector3.Normalize(Vector3.Sub(vectorB, vectorA));
end

Vector3.Clamp = function(vector, maxLength)
	local length = Vector3.Length(vector);

	if (length > maxLength)
		return Vector3.Mul(Vector3.Normalized(vector), maxLength);
	end

	return Vector3.Copy(vector);
end

Vector3.Rotate = function(vector, aX, aY, aZ)
	local sineX = sin(aX);
	local cosineX = cos(aX);

	local sineY = sin(aY);
	local cosineY = cos(aY);

	local sineZ = sin(aZ);
	local cosineZ = cos(aZ);

	local newVector = Vector3.New(
		vector.x,
		vector.y,
		vector.z
	);

	-- X rotation
	newVector.y = FixedMul(newVector.y, cosineY) - FixedMul(newVector.z, sineY);
	newVector.z = FixedMul(newVector.y, sineY) + FixedMul(newVector.z, cosineY);

	-- Y rotation
	newVector.x = FixedMul(newVector.x, cosineX) + FixedMul(newVector.z, sineX);
	newVector.z = -FixedMul(newVector.x, sineX) + FixedMul(newVector.z, cosineX);

	-- Z rotation
	newVector.x = FixedMul(newVector.x, cosineZ) - FixedMul(newVector.y, sineZ);
	newVector.y = FixedMul(newVector.x, sineZ) + FixedMul(newVector.y, cosineZ);

	return newVector;
end

Vector3.Dot = function(vectorA, vectorB)
	return FixedMul(vectorA.x, vectorB.x) + FixedMul(vectorA.y, vectorB.y) + FixedMul(vectorA.z, vectorB.z);
end

Vector3.Cross = function(vectorA, vectorB)
	return Vector3.New(
		FixedMul(vectorA.y, vectorB.z) - FixedMul(vectorA.z, vectorB.y),
		FixedMul(vectorA.z, vectorB.x) - FixedMul(vectorA.x, vectorB.z),
		FixedMul(vectorA.x, vectorB.y) - FixedMul(vectorA.y, vectorB.x)
	);
end

Vector3.MobjPositionVector = function(mobj)
	return Vector3.New(mobj.x, mobj.y, mobj.z);
end

Vector3.MobjMomentumVector = function(mobj)
	return Vector3.New(mobj.momx, mobj.momy, mobj.momz);
end

Vector3.SetMobjMomentum = function(mobj, vector)
	if not (mobj and mobj.valid)
		return;
	end

	mobj.momx = vector.x;
	mobj.momy = vector.y;
	mobj.momz = vector.z;
end

Vector3.Meta = {
	__add = Vector3.Add,
	__sub = Vector3.Sub,
	__mul = Vector3.IntegerMul,
	__div = Vector3.IntegerDiv,
	__unm = Vector3.Negate,
	__eq = Vector3.Equal,
};
registerMetatable(Vector3.Meta);

rawset(_G, "Vector3", Vector3);
