diff --git a/include/rest_rpc/codec.h b/include/rest_rpc/codec.h index 4905353..d28c9b1 100644 --- a/include/rest_rpc/codec.h +++ b/include/rest_rpc/codec.h @@ -5,18 +5,44 @@ #include #include -namespace rest_rpc { +namespace user_codec { +struct rest_adl_tag {}; +} // namespace user_codec -struct msgpack_codec { +namespace rest_rpc { +namespace detail { +template +struct has_user_pack : std::false_type {}; + +template +struct has_user_pack< + std::void_t< + // AdlTag{} trigger ADL lookup user_codec namespace + decltype(serialize(std::declval(), + std::declval()...))>, + Args...> : std::true_type {}; + +template +inline constexpr bool has_user_pack_v = has_user_pack::value; + +} // namespace detail + +struct rpc_codec { template inline static auto pack_args(Args &&...args) { if constexpr (sizeof...(Args) == 0) { return std::string_view{}; } else if constexpr (sizeof...(Args) == 1 && util::is_basic_v) { return pack_one(std::forward(args)...); } else { - msgpack::sbuffer buffer(2 * 1024); - msgpack::pack(buffer, std::forward_as_tuple(std::forward(args)...)); - return std::string(buffer.data(), buffer.size()); + if constexpr (detail::has_user_pack_v) { + return serialize(user_codec::rest_adl_tag{}, + std::forward(args)...); + } else { + msgpack::sbuffer buffer(2 * 1024); + msgpack::pack(buffer, + std::forward_as_tuple(std::forward(args)...)); + return std::string(buffer.data(), buffer.size()); + } } } @@ -33,12 +59,16 @@ struct msgpack_codec { } else if constexpr (std::is_same_v) { return data; } else { - try { - static msgpack::unpacked msg; - msgpack::unpack(msg, data.data(), data.size()); - return msg.get().as(); - } catch (...) { - throw std::invalid_argument("unpack failed: Args not match!"); + if constexpr (detail::has_user_pack_v) { + return deserialize(user_codec::rest_adl_tag{}, data); + } else { + try { + static msgpack::unpacked msg; + msgpack::unpack(msg, data.data(), data.size()); + return msg.get().as(); + } catch (...) { + throw std::invalid_argument("unpack failed: Args not match!"); + } } } } diff --git a/include/rest_rpc/rpc_client.hpp b/include/rest_rpc/rpc_client.hpp index 739981b..2808de7 100644 --- a/include/rest_rpc/rpc_client.hpp +++ b/include/rest_rpc/rpc_client.hpp @@ -164,7 +164,7 @@ private: template asio::awaitable> call_impl(rest_rpc_header &header, Args &&...args) { - auto buf = msgpack_codec::pack_args(std::forward(args)...); + auto buf = rpc_codec::pack_args(std::forward(args)...); header.body_len = buf.size(); if (cross_ending_) { prepare_for_send(header); @@ -227,10 +227,10 @@ private: result.ec = (rpc_errc)socket_->body_[0]; if constexpr (!std::is_void_v) { if constexpr (util::is_basic_v) { - result.value = msgpack_codec::unpack(std::string_view( + result.value = rpc_codec::unpack(std::string_view( socket_->body_.data() + 1, resp_header.body_len - 1)); } else { - auto tp = msgpack_codec::unpack>(std::string_view( + auto tp = rpc_codec::unpack>(std::string_view( socket_->body_.data() + 1, resp_header.body_len - 1)); result.value = std::move(std::get<0>(tp)); } diff --git a/include/rest_rpc/rpc_connection.hpp b/include/rest_rpc/rpc_connection.hpp index 07bc3d3..633254a 100644 --- a/include/rest_rpc/rpc_connection.hpp +++ b/include/rest_rpc/rpc_connection.hpp @@ -238,7 +238,7 @@ asio::awaitable rpc_context::response(Args &&...args) { co_return make_error_code(rpc_errc::rpc_context_init_failed); } - rpc_result result(msgpack_codec::pack_args(std::forward(args)...)); + rpc_result result(rpc_codec::pack_args(std::forward(args)...)); has_response_ = true; co_return co_await conn_->response(result); } diff --git a/include/rest_rpc/rpc_router.hpp b/include/rest_rpc/rpc_router.hpp index 8b8628b..1e25a66 100644 --- a/include/rest_rpc/rpc_router.hpp +++ b/include/rest_rpc/rpc_router.hpp @@ -148,15 +148,15 @@ private: } else { if constexpr (std::is_void_v) { if constexpr (is_awaitable_v) { - ret = msgpack_codec::pack_args(co_await f()); + ret = rpc_codec::pack_args(co_await f()); } else { - ret = msgpack_codec::pack_args(f()); + ret = rpc_codec::pack_args(f()); } } else { if constexpr (is_awaitable_v) { - ret = msgpack_codec::pack_args(co_await (*self.*f)()); + ret = rpc_codec::pack_args(co_await (*self.*f)()); } else { - ret = msgpack_codec::pack_args((*self.*f)()); + ret = rpc_codec::pack_args((*self.*f)()); } } } @@ -169,32 +169,30 @@ private: if constexpr (is_void_v) { if constexpr (std::is_void_v) { if constexpr (is_awaitable_v) { - co_await f(msgpack_codec::unpack(str)); + co_await f(rpc_codec::unpack(str)); } else { - f(msgpack_codec::unpack(str)); + f(rpc_codec::unpack(str)); } } else { if constexpr (is_awaitable_v) { - co_await (*self.*f)(msgpack_codec::unpack(str)); + co_await (*self.*f)(rpc_codec::unpack(str)); } else { - (*self.*f)(msgpack_codec::unpack(str)); + (*self.*f)(rpc_codec::unpack(str)); } } } else { if constexpr (std::is_void_v) { if constexpr (is_awaitable_v) { - ret = msgpack_codec::pack_args( - co_await f(msgpack_codec::unpack(str))); + ret = rpc_codec::pack_args(co_await f(rpc_codec::unpack(str))); } else { - ret = msgpack_codec::pack_args(f(msgpack_codec::unpack(str))); + ret = rpc_codec::pack_args(f(rpc_codec::unpack(str))); } } else { if constexpr (is_awaitable_v) { - ret = msgpack_codec::pack_args( - co_await (*self.*f)(msgpack_codec::unpack(str))); + ret = rpc_codec::pack_args( + co_await (*self.*f)(rpc_codec::unpack(str))); } else { - ret = msgpack_codec::pack_args( - (*self.*f)(msgpack_codec::unpack(str))); + ret = rpc_codec::pack_args((*self.*f)(rpc_codec::unpack(str))); } } } @@ -204,7 +202,7 @@ private: template asio::awaitable handle_more_args(std::string_view str, const F &f, rpc_result &ret, Self *self) { - auto tp = msgpack_codec::unpack(str); + auto tp = rpc_codec::unpack(str); if constexpr (std::is_void_v) { if constexpr (std::is_void_v) { if constexpr (is_awaitable_v) { @@ -230,19 +228,19 @@ private: } else { if constexpr (std::is_void_v) { if constexpr (is_awaitable_v) { - ret = msgpack_codec::pack_args(co_await std::apply(f, tp)); + ret = rpc_codec::pack_args(co_await std::apply(f, tp)); } else { - ret = msgpack_codec::pack_args(std::apply(f, tp)); + ret = rpc_codec::pack_args(std::apply(f, tp)); } } else { if constexpr (is_awaitable_v) { - ret = msgpack_codec::pack_args(co_await std::apply( + ret = rpc_codec::pack_args(co_await std::apply( [self, &f](auto &&...args) { return (*self.*f)(std::forward(args)...); }, tp)); } else { - ret = msgpack_codec::pack_args(std::apply( + ret = rpc_codec::pack_args(std::apply( [self, &f](auto &&...args) { return (*self.*f)(std::forward(args)...); }, diff --git a/include/rest_rpc/rpc_server.hpp b/include/rest_rpc/rpc_server.hpp index 9664364..87a10d1 100644 --- a/include/rest_rpc/rpc_server.hpp +++ b/include/rest_rpc/rpc_server.hpp @@ -113,8 +113,7 @@ public: auto conns = get_connections(); for (auto &[_, conn] : conns) { if (conn->topic_id() == id) { - co_await conn->response(msgpack_codec::pack_args(std::forward(t)), - id); + co_await conn->response(rpc_codec::pack_args(std::forward(t)), id); } } } diff --git a/tests/test_rest_rpc.cpp b/tests/test_rest_rpc.cpp index 4fac2a5..d941216 100644 --- a/tests/test_rest_rpc.cpp +++ b/tests/test_rest_rpc.cpp @@ -243,8 +243,8 @@ asio::awaitable test_router() { } { - auto s = msgpack_codec::pack_args(1); - auto s1 = msgpack_codec::pack_args("test"); + auto s = rpc_codec::pack_args(1); + auto s1 = rpc_codec::pack_args("test"); auto r = co_await router.route(get_key(), s); auto r1 = co_await router.route(get_key(), s1); @@ -253,7 +253,7 @@ asio::awaitable test_router() { std::cout << "\n"; } - auto args = msgpack_codec::pack_args(1, 2); + auto args = rpc_codec::pack_args(1, 2); std::string_view str(args.data(), args.size()); { @@ -262,12 +262,12 @@ asio::awaitable test_router() { std::cout << "\n"; } - auto args1 = msgpack_codec::pack_args("it is a test"); + auto args1 = rpc_codec::pack_args("it is a test"); std::string_view str1(args1.data(), args1.size()); { auto result = co_await router.route(get_key<&dummy::add>(), str); - auto r = msgpack_codec::unpack(result.result); + auto r = rpc_codec::unpack(result.result); auto result1 = co_await router.route(get_key<&dummy::foo>(), str1); CHECK(r == 3); CHECK(result1.ec == rpc_errc::ok); @@ -293,10 +293,10 @@ TEST_CASE("test rpc_connection") { person p{1, "tom", 20}; - auto buf = msgpack_codec::pack_args(p); + auto buf = rpc_codec::pack_args(p); auto ret = sync_wait(get_global_executor(), router.route(get_key(), buf)); - auto tp = msgpack_codec::unpack>(ret.data()); + auto tp = rpc_codec::unpack>(ret.data()); dummy d{}; router.register_handler<&dummy::add>(&d); auto conn = std::make_shared(std::move(socket), conn_id, @@ -334,12 +334,12 @@ TEST_CASE("test server start") { static_assert(util::CharArrayRef); static_assert(util::CharArray); - msgpack_codec::pack_args(); - msgpack_codec::pack_args(1, 2); - auto s1 = msgpack_codec::pack_args("test"); - auto s2 = msgpack_codec::pack_args(std::string_view("test2")); - auto s3 = msgpack_codec::pack_args(std::string("test2")); - auto s5 = msgpack_codec::pack_args(123); + rpc_codec::pack_args(); + rpc_codec::pack_args(1, 2); + auto s1 = rpc_codec::pack_args("test"); + auto s2 = rpc_codec::pack_args(std::string_view("test2")); + auto s3 = rpc_codec::pack_args(std::string("test2")); + auto s5 = rpc_codec::pack_args(123); // auto future = asio::co_spawn(cl.get_executor(), // cl.connect("127.0.0.1:9005"), asio::use_future); auto conn_ec = @@ -555,6 +555,40 @@ TEST_CASE("test server address") { thd.join(); } +bool in_user_pack = false; +bool in_user_unpack = false; +namespace user_codec { +// adl lookup in user_codec namespace +template +std::string serialize(rest_adl_tag, Args &&...args) { + in_user_pack = true; + msgpack::sbuffer buffer(2 * 1024); + msgpack::pack(buffer, std::forward_as_tuple(std::forward(args)...)); + return std::string(buffer.data(), buffer.size()); +} + +template T deserialize(rest_adl_tag, std::string_view data) { + try { + in_user_unpack = true; + static msgpack::unpacked msg; + msgpack::unpack(msg, data.data(), data.size()); + return msg.get().as(); + } catch (...) { + return T{}; + } +} +} // namespace user_codec + +TEST_CASE("test user codec") { + auto buf = rpc_codec::pack_args( + std::make_tuple(1, "tom", 20)); + CHECK(in_user_pack); + + std::string_view str(buf.data(), buf.size()); + rpc_codec::unpack>(str); + CHECK(in_user_unpack); +} + // doctest comments // 'function' : must be 'attribute' - see issue #182 DOCTEST_MSVC_SUPPRESS_WARNING_WITH_PUSH(4007)