Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 45 additions & 1 deletion Source/LuaBridge/detail/CFunctions.h
Original file line number Diff line number Diff line change
Expand Up @@ -254,6 +254,37 @@ inline bool is_metamethod(std::string_view method_name)
return result != metamethods.end() && *result == method_name;
}

/**
* @brief Check if a method name is a binary operator metamethod.
*/
inline bool is_binary_operator_metamethod(std::string_view method_name)
{
static constexpr auto metamethods = make_array<std::string_view>(
"__add",
"__band",
"__bor",
"__bxor",
"__concat",
"__div",
"__eq",
"__idiv",
"__le",
"__lt",
"__mod",
"__mul",
"__pow",
"__shl",
"__shr",
"__sub"
);

if (method_name.size() <= 2 || method_name[0] != '_' || method_name[1] != '_')
return false;

auto result = std::lower_bound(metamethods.begin(), metamethods.end(), method_name);
return result != metamethods.end() && *result == method_name;
}

inline void rawset_super_method(lua_State* L, int tableIndex, const char* key)
{
LUABRIDGE_ASSERT(key != nullptr);
Expand Down Expand Up @@ -2251,6 +2282,18 @@ bool overload_type_checker(lua_State* L, int start)
return overload_check_args<ArgsPack>(L, start);
}

/**
* @brief Type checker for reversed proxy functions (binary operators with swapped operands).
*
* Checks all arguments from absolute stack index 1, ignoring @p start, because the class object
* is not the first operand on the stack (eg. `3 * vec`).
*/
template <class ArgsPack>
bool reversed_overload_type_checker(lua_State* L, int)
{
return overload_check_args<ArgsPack>(L, 1);
}

//=================================================================================================
/**
* @brief lua_CFunction to resolve an invocation between several overloads.
Expand Down Expand Up @@ -2479,7 +2522,8 @@ template <class T, class F, class = std::enable_if<
!std::is_member_function_pointer_v<F>>>
void push_member_function(lua_State* L, F&& f, const char* debugname)
{
static_assert(std::is_same_v<T, remove_cvref_t<std::remove_pointer_t<function_argument_or_void_t<0, F>>>>);
static_assert(std::is_same_v<T, remove_cvref_t<std::remove_pointer_t<function_argument_or_void_t<0, F>>>> ||
is_reversed_proxy_function_v<T, F>);

lua_newuserdata_aligned<F>(L, std::forward<F>(f));
lua_pushcclosure_x(L, &invoke_proxy_functor<F>, debugname, 1);
Expand Down
25 changes: 24 additions & 1 deletion Source/LuaBridge/detail/FuncTraits.h
Original file line number Diff line number Diff line change
Expand Up @@ -541,6 +541,28 @@ inline static constexpr bool is_const_proxy_function_v =
is_proxy_member_function_v<T, F> &&
std::is_const_v<std::remove_reference_t<std::remove_pointer_t<function_argument_or_void_t<0, F>>>>;

//=================================================================================================
/**
* @brief A constexpr check for reversed proxy member functions (binary operators).
*
* A reversed proxy function is a binary callable where T is the second argument instead of the
* first, as needed by binary operator metamethods invoked with swapped operands (eg. `3 * vec`).
*
* @tparam T Type where the callable should be able to operate.
* @tparam F Callable object.
*/
template <class T, class F>
inline static constexpr bool is_reversed_proxy_function_v =
!std::is_member_function_pointer_v<F> &&
!is_proxy_member_function_v<T, F> &&
function_arity_v<F> == 2 &&
std::is_same_v<T, remove_cvref_t<std::remove_pointer_t<function_argument_or_void_t<1, F>>>>;

template <class T, class F>
inline static constexpr bool is_const_reversed_proxy_function_v =
is_reversed_proxy_function_v<T, F> &&
std::is_const_v<std::remove_reference_t<std::remove_pointer_t<function_argument_or_void_t<1, F>>>>;

//=================================================================================================
/**
* @brief An integral constant expression that gives the number of arguments excluding one type (usually used with lua_State*) accepted by the callable object.
Expand Down Expand Up @@ -593,7 +615,8 @@ inline static constexpr std::size_t member_function_arity_excluding_v = member_f
template <class T, class F>
static constexpr bool is_const_function =
detail::is_const_member_function_pointer_v<F> ||
(detail::function_arity_v<F> > 0 && detail::is_const_proxy_function_v<T, F>);
(detail::function_arity_v<F> > 0 && detail::is_const_proxy_function_v<T, F>) ||
detail::is_const_reversed_proxy_function_v<T, F>;

template <class T, class... Fs>
inline static constexpr std::size_t const_functions_count = (0 + ... + (is_const_function<T, Fs> ? 1 : 0));
Expand Down
21 changes: 21 additions & 0 deletions Source/LuaBridge/detail/Namespace.h
Original file line number Diff line number Diff line change
Expand Up @@ -950,6 +950,15 @@ class Namespace : public detail::Registrar
#endif
}

if constexpr ((detail::is_reversed_proxy_function_v<T, Functions> || ...))
{
if (!detail::is_binary_operator_metamethod(name))
{
throw_or_assert<std::logic_error>("reversed functions are only allowed on binary operator metamethods");
return *this;
}
}

if constexpr (sizeof...(Functions) == 1)
{
([&]
Expand Down Expand Up @@ -995,6 +1004,12 @@ class Namespace : public detail::Registrar
entry.arity = static_cast<int>(detail::member_function_arity_excluding_v<T, Functions, lua_State*>);
entry.checker = &detail::overload_type_checker<ArgsPack>;
}
else if constexpr (detail::is_reversed_proxy_function_v<T, Functions>)
{
using ArgsPack = detail::function_arguments_t<Functions>;
entry.arity = static_cast<int>(detail::member_function_arity_excluding_v<T, Functions, lua_State*>) - 1;
entry.checker = &detail::reversed_overload_type_checker<ArgsPack>;
}
else
{
using ArgsPack = detail::function_arguments_t<Functions>;
Expand Down Expand Up @@ -1052,6 +1067,12 @@ class Namespace : public detail::Registrar
entry.arity = static_cast<int>(detail::member_function_arity_excluding_v<T, Functions, lua_State*>);
entry.checker = &detail::overload_type_checker<ArgsPack>;
}
else if constexpr (detail::is_reversed_proxy_function_v<T, Functions>)
{
using ArgsPack = detail::function_arguments_t<Functions>;
entry.arity = static_cast<int>(detail::member_function_arity_excluding_v<T, Functions, lua_State*>) - 1;
entry.checker = &detail::reversed_overload_type_checker<ArgsPack>;
}
else
{
using ArgsPack = detail::function_arguments_t<Functions>;
Expand Down
195 changes: 195 additions & 0 deletions Tests/Source/ClassTests.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -169,6 +169,67 @@ class Bar
Foo<T> foo;
};

class Vec2
{
public:
Vec2(float x, float y)
: x(x), y(y)
{
}

friend Vec2 operator*(const Vec2& vec, float scalar)
{
return Vec2(vec.x * scalar, vec.y * scalar);
}

friend Vec2 operator*(float scalar, const Vec2& vec)
{
return Vec2(vec.x * scalar, vec.y * scalar);
}

friend Vec2 operator+(const Vec2& vec, float scalar)
{
return Vec2(vec.x + scalar, vec.y + scalar);
}

friend Vec2 operator+(float scalar, const Vec2& vec)
{
return Vec2(scalar + vec.x, scalar + vec.y);
}

friend Vec2 operator-(const Vec2& vec, float scalar)
{
return Vec2(vec.x - scalar, vec.y - scalar);
}

friend Vec2 operator-(float scalar, const Vec2& vec)
{
return Vec2(scalar - vec.x, scalar - vec.y);
}

friend Vec2 operator/(const Vec2& vec, float scalar)
{
return Vec2(vec.x / scalar, vec.y / scalar);
}

friend Vec2 operator/(float scalar, const Vec2& vec)
{
return Vec2(scalar / vec.x, scalar / vec.y);
}

friend bool operator==(const Vec2& lhs, const Vec2& rhs)
{
return lhs.x == rhs.x && lhs.y == rhs.y;
}

float getX() const { return x; }
float getY() const { return y; }

private:
float x;
float y;
};

} // namespace

TEST_F(ClassTests, Assignment)
Expand Down Expand Up @@ -2572,6 +2633,122 @@ TEST_F(ClassMetaMethods, __mul)
ASSERT_EQ(10, result<Int>().data);
}

TEST_F(ClassMetaMethods, ReversedAddOperator)
{
luabridge::getGlobalNamespace(L)
.beginClass<Vec2>("Vec2")
.addConstructor<void (*)(float, float)>()
.addFunction("__add",
[](const Vec2& vec, float scalar) { return vec + scalar; },
[](float scalar, const Vec2& vec) { return scalar + vec; })
.endClass();

runLua("result = Vec2 (1, 2) + 3");
ASSERT_TRUE(result().isUserdata());
ASSERT_EQ(Vec2(4, 5), result<Vec2>());

runLua("result = 3 + Vec2 (1, 2)");
ASSERT_TRUE(result().isUserdata());
ASSERT_EQ(Vec2(4, 5), result<Vec2>());
}

TEST_F(ClassMetaMethods, ReversedSubOperator)
{
luabridge::getGlobalNamespace(L)
.beginClass<Vec2>("Vec2")
.addConstructor<void (*)(float, float)>()
.addFunction("__sub",
[](const Vec2& vec, float scalar) { return vec - scalar; },
[](float scalar, const Vec2& vec) { return scalar - vec; })
.endClass();

runLua("result = Vec2 (5, 7) - 2");
ASSERT_TRUE(result().isUserdata());
ASSERT_EQ(Vec2(3, 5), result<Vec2>());

runLua("result = 10 - Vec2 (1, 2)");
ASSERT_TRUE(result().isUserdata());
ASSERT_EQ(Vec2(9, 8), result<Vec2>());
}

TEST_F(ClassMetaMethods, ReversedMulOperator)
{
luabridge::getGlobalNamespace(L)
.beginClass<Vec2>("Vec2")
.addConstructor<void (*)(float, float)>()
.addFunction("__mul",
[](const Vec2& vec, float scalar) { return vec * scalar; },
[](float scalar, const Vec2& vec) { return scalar * vec; })
.endClass();

runLua("result = Vec2 (1, 2) * 3");
ASSERT_TRUE(result().isUserdata());
ASSERT_EQ(Vec2(3, 6), result<Vec2>());

runLua("result = 3 * Vec2 (1, 2)");
ASSERT_TRUE(result().isUserdata());
ASSERT_EQ(Vec2(3, 6), result<Vec2>());
}

TEST_F(ClassMetaMethods, ReversedDivOperator)
{
luabridge::getGlobalNamespace(L)
.beginClass<Vec2>("Vec2")
.addConstructor<void (*)(float, float)>()
.addFunction("__div",
[](const Vec2& vec, float scalar) { return vec / scalar; },
[](float scalar, const Vec2& vec) { return scalar / vec; })
.endClass();

runLua("result = Vec2 (6, 8) / 2");
ASSERT_TRUE(result().isUserdata());
ASSERT_EQ(Vec2(3, 4), result<Vec2>());

runLua("result = 12 / Vec2 (2, 4)");
ASSERT_TRUE(result().isUserdata());
ASSERT_EQ(Vec2(6, 3), result<Vec2>());
}

TEST_F(ClassMetaMethods, ReversedOperatorOnConstInstance)
{
luabridge::getGlobalNamespace(L)
.beginClass<Vec2>("Vec2")
.addConstructor<void (*)(float, float)>()
.addFunction("__mul",
[](const Vec2& vec, float scalar) { return vec * scalar; },
[](float scalar, const Vec2& vec) { return scalar * vec; })
.endClass();

const Vec2 vec(1, 2);
luabridge::setGlobal(L, &vec, "v");

runLua("result = v * 3");
ASSERT_TRUE(result().isUserdata());
ASSERT_EQ(Vec2(3, 6), result<Vec2>());

runLua("result = 3 * v");
ASSERT_TRUE(result().isUserdata());
ASSERT_EQ(Vec2(3, 6), result<Vec2>());
}

TEST_F(ClassMetaMethods, ReversedOperatorOnly)
{
luabridge::getGlobalNamespace(L)
.beginClass<Vec2>("Vec2")
.addConstructor<void (*)(float, float)>()
.addFunction("__mul",
[](float scalar, const Vec2& vec) { return scalar * vec; })
.endClass();

runLua("result = 3 * Vec2 (1, 2)");
ASSERT_TRUE(result().isUserdata());
ASSERT_EQ(Vec2(3, 6), result<Vec2>());

#if LUABRIDGE_HAS_EXCEPTIONS
ASSERT_ANY_THROW(runLua("result = Vec2 (1, 2) * 3"));
#endif
}

TEST_F(ClassMetaMethods, __div)
{
typedef Class<int, EmptyBase> Int;
Expand Down Expand Up @@ -2776,6 +2953,24 @@ TEST_F(ClassMetaMethods, __gcForbidden)
.endClass(),
std::exception);
}

TEST_F(ClassMetaMethods, ReversedFunctionsForbiddenOutsideBinaryOperators)
{
ASSERT_THROW(luabridge::getGlobalNamespace(L)
.beginClass<Vec2>("Vec2")
.addFunction("mul",
[](const Vec2& vec, float scalar) { return vec * scalar; },
[](float scalar, const Vec2& vec) { return scalar * vec; })
.endClass(),
std::exception);

ASSERT_THROW(luabridge::getGlobalNamespace(L)
.beginClass<Vec2>("Vec2")
.addFunction("__tostring",
[](float scalar, const Vec2& vec) { return scalar * vec; })
.endClass(),
std::exception);
}
#endif

TEST_F(ClassMetaMethods, MetamethodsShouldNotBePartOfClassInstances)
Expand Down
Loading