From 25b0a445b96a9e90d0ada7f3584dd54f190d0213 Mon Sep 17 00:00:00 2001 From: Tyler Neely Date: Fri, 8 Jul 2022 08:21:37 +0000 Subject: [PATCH] Properly fill response promises in the simulator --- src/io/v3/simulator_handle.hpp | 29 +++++++++++++++++------------ tests/simulation/basic_request.cpp | 2 +- 2 files changed, 18 insertions(+), 13 deletions(-) diff --git a/src/io/v3/simulator_handle.hpp b/src/io/v3/simulator_handle.hpp index 1a47e295f..b0f001900 100644 --- a/src/io/v3/simulator_handle.hpp +++ b/src/io/v3/simulator_handle.hpp @@ -135,13 +135,15 @@ class OpaquePromise { ptr_((void *)promise.release()), dtor_([](void *ptr) { static_cast *>(ptr)->~ResponsePromise(); }), is_awaited_([](void *ptr) { return static_cast *>(ptr)->IsAwaited(); }), - fill_([](void *this_ptr, std::any msg_any) { + fill_([](void *this_ptr, OpaqueMessage opaque_message) { std::cout << "expecting typeid " << typeid(T).name() << std::endl; - std::cout << "got typeid " << msg_any.type().name() << std::endl; - MG_ASSERT(typeid(T) == msg_any.type(), "type id mismatch"); - ResponseResult message = std::any_cast>(std::move(msg_any)); - auto promise = static_cast *>(this_ptr); - promise->Fill(std::move(message)); + std::cout << "got typeid " << opaque_message.message.type().name() << std::endl; + T message = std::any_cast(std::move(opaque_message.message)); + auto response_envelope = ResponseEnvelope{.message = std::move(message), + .request_id = opaque_message.request_id, + .from_address = opaque_message.from_address}; + ResponsePromise *promise = static_cast *>(this_ptr); + promise->Fill(std::move(response_envelope)); }), time_out_([](void *ptr) { ResponseResult result = TimedOut{}; @@ -160,7 +162,7 @@ class OpaquePromise { void Fill(OpaqueMessage &&opaque_message) { MG_ASSERT(ptr_ != nullptr); - fill_(ptr_, std::move(opaque_message.message)); + fill_(ptr_, std::move(opaque_message)); } ~OpaquePromise() { @@ -174,7 +176,7 @@ class OpaquePromise { void *ptr_; std::function dtor_; std::function is_awaited_; - std::function fill_; + std::function fill_; std::function time_out_; }; @@ -204,6 +206,9 @@ class SimulatorHandle { } } + std::cout << "wait count: " << blocked_servers << std::endl; + std::cout << "srv count: " << servers_ << std::endl; + if (blocked_servers < servers_) { // we only need to advance the simulator when all // servers have reached a quiescent state, blocked @@ -240,9 +245,6 @@ class SimulatorHandle { om_vec->second.emplace_back(std::move(opaque_message)); } - std::cout << "wait count: " << blocked_servers << std::endl; - std::cout << "srv count: " << servers_ << std::endl; - cv_.notify_all(); return true; @@ -314,7 +316,10 @@ class SimulatorHandle { template void Send(Address to_address, Address from_address, uint64_t request_id, M message) { - std::abort(); + std::unique_lock lock(mu_); + std::any message_any(std::move(message)); + OpaqueMessage om{.from_address = from_address, .request_id = request_id, .message = std::move(message_any)}; + in_flight_.emplace_back(std::make_pair(std::move(to_address), std::move(om))); } private: diff --git a/tests/simulation/basic_request.cpp b/tests/simulation/basic_request.cpp index 093839a25..13dd45eca 100644 --- a/tests/simulation/basic_request.cpp +++ b/tests/simulation/basic_request.cpp @@ -39,7 +39,7 @@ int main() { auto srv_addr = Address::TestAddress(2); Io cli_io = simulator.Register(cli_addr, false); - Io srv_io = simulator.Register(srv_addr, true); + Io srv_io = simulator.Register(srv_addr, false); // send request RequestMsg cli_req;