diff --git a/src/spider/client/Driver.hpp b/src/spider/client/Driver.hpp index 558cfae83..5baafe08f 100644 --- a/src/spider/client/Driver.hpp +++ b/src/spider/client/Driver.hpp @@ -228,6 +228,8 @@ class Driver { return Job{ job_id, + Job::JobSource::Driver, + m_id, m_metadata_storage, m_data_storage, m_storage_factory, @@ -291,6 +293,8 @@ class Driver { return Job{ job_id, + Job::JobSource::Driver, + m_id, m_metadata_storage, m_data_storage, m_storage_factory, diff --git a/src/spider/client/Job.hpp b/src/spider/client/Job.hpp index c79b6d668..58ffa5083 100644 --- a/src/spider/client/Job.hpp +++ b/src/spider/client/Job.hpp @@ -157,21 +157,34 @@ class Job { } private: - Job(boost::uuids::uuid id, + enum class JobSource : uint8_t { + Driver, + Task, + }; + + Job(boost::uuids::uuid const id, + JobSource const source, + boost::uuids::uuid const source_id, std::shared_ptr metadata_storage, std::shared_ptr data_storage, std::shared_ptr storage_factory) : m_id{id}, + m_source{source}, + m_source_id{source_id}, m_metadata_storage{std::move(metadata_storage)}, m_data_storage{std::move(data_storage)}, m_storage_factory{std::move(storage_factory)} {} - Job(boost::uuids::uuid id, + Job(boost::uuids::uuid const id, + JobSource const source, + boost::uuids::uuid const source_id, std::shared_ptr metadata_storage, std::shared_ptr data_storage, std::shared_ptr storage_factory, std::shared_ptr conn) : m_id{id}, + m_source{source}, + m_source_id{source_id}, m_metadata_storage{std::move(metadata_storage)}, m_data_storage{std::move(data_storage)}, m_storage_factory{std::move(storage_factory)}, @@ -242,7 +255,21 @@ class Job { if (!optional_data_id.has_value()) { throw ConnectionException{fmt::format("Output data ID is missing")}; } - err = m_data_storage->get_data(conn, optional_data_id.value(), &data); + if (m_source == JobSource::Driver) { + err = m_data_storage->get_driver_data( + conn, + m_source_id, + optional_data_id.value(), + &data + ); + } else { + err = m_data_storage->get_task_data( + conn, + m_source_id, + optional_data_id.value(), + &data + ); + } if (!err.success()) { throw ConnectionException{ fmt::format("Failed to get data: {}", err.description) @@ -302,7 +329,21 @@ class Job { if (!optional_data_id.has_value()) { throw ConnectionException{fmt::format("Output data ID is missing")}; } - err = m_data_storage->get_data(conn, optional_data_id.value(), &data); + if (m_source == JobSource::Driver) { + err = m_data_storage->get_driver_data( + conn, + m_source_id, + optional_data_id.value(), + &data + ); + } else { + err = m_data_storage->get_task_data( + conn, + m_source_id, + optional_data_id.value(), + &data + ); + } if (!err.success()) { throw ConnectionException{ fmt::format("Failed to get data: {}", err.description) @@ -337,6 +378,8 @@ class Job { // NOLINTEND(readability-function-cognitive-complexity) boost::uuids::uuid m_id; + JobSource m_source; + boost::uuids::uuid m_source_id; std::shared_ptr m_metadata_storage; std::shared_ptr m_data_storage; std::shared_ptr m_storage_factory; diff --git a/src/spider/client/TaskContext.hpp b/src/spider/client/TaskContext.hpp index db3b9292d..904a944f7 100644 --- a/src/spider/client/TaskContext.hpp +++ b/src/spider/client/TaskContext.hpp @@ -172,7 +172,14 @@ class TaskContext { throw ConnectionException(fmt::format("Failed to start job: {}", err.description)); } - return Job{job_id, m_metadata_store, m_data_store, m_storage_factory}; + return Job{ + job_id, + Job::JobSource::Task, + m_task_id, + m_metadata_store, + m_data_store, + m_storage_factory + }; } /** @@ -224,7 +231,14 @@ class TaskContext { throw ConnectionException(fmt::format("Failed to start job: {}", err.description)); } - return Job{job_id, m_metadata_store, m_data_store, m_storage_factory}; + return Job{ + job_id, + Job::JobSource::Task, + m_task_id, + m_metadata_store, + m_data_store, + m_storage_factory + }; } /** diff --git a/src/spider/storage/DataStorage.hpp b/src/spider/storage/DataStorage.hpp index 041f04871..79713111e 100644 --- a/src/spider/storage/DataStorage.hpp +++ b/src/spider/storage/DataStorage.hpp @@ -32,6 +32,38 @@ class DataStorage { = 0; virtual auto get_data(StorageConnection& conn, boost::uuids::uuid id, Data* data) -> StorageErr = 0; + /** + * Get a data object and register a reference for it from the given driver in a single + * transaction. + * @param conn + * @param driver_id + * @param data_id + * @param data output data + * @return StorageErr::Success if the transaction succeed. Error types otherwise. + */ + virtual auto get_driver_data( + StorageConnection& conn, + boost::uuids::uuid driver_id, + boost::uuids::uuid data_id, + Data* data + ) -> StorageErr + = 0; + /** + * Get a data object and register a reference for it from the given task in a single + * transaction. + * @param conn + * @param task_id + * @param data_id + * @param data output data + * @return StorageErr::Success if the transaction succeed. Error types otherwise. + */ + virtual auto get_task_data( + StorageConnection& conn, + boost::uuids::uuid task_id, + boost::uuids::uuid data_id, + Data* data + ) -> StorageErr + = 0; virtual auto set_data_locality(StorageConnection& conn, Data const& data) -> StorageErr = 0; virtual auto remove_data(StorageConnection& conn, boost::uuids::uuid id) -> StorageErr = 0; virtual auto diff --git a/src/spider/storage/mysql/MySqlStorage.cpp b/src/spider/storage/mysql/MySqlStorage.cpp index 79df41e91..7678f69b8 100644 --- a/src/spider/storage/mysql/MySqlStorage.cpp +++ b/src/spider/storage/mysql/MySqlStorage.cpp @@ -2074,42 +2074,114 @@ auto MySqlDataStorage::add_task_data( return StorageErr{}; } -auto MySqlDataStorage::get_data(StorageConnection& conn, boost::uuids::uuid id, Data* data) +auto MySqlDataStorage::get_data_with_locality( + StorageConnection& conn, + boost::uuids::uuid const id, + Data* data +) -> StorageErr { + std::unique_ptr statement( + static_cast(conn)->prepareStatement( + "SELECT `id`, `value`, `hard_locality` FROM `data` WHERE `id` = ?" + ) + ); + sql::bytes id_bytes = uuid_get_bytes(id); + statement->setBytes(1, &id_bytes); + std::unique_ptr res(statement->executeQuery()); + if (res->rowsCount() == 0) { + static_cast(conn)->rollback(); + return StorageErr{ + StorageErrType::KeyNotFoundErr, + fmt::format("no data with id {}", boost::uuids::to_string(id)) + }; + } + res->next(); + *data = Data{id, get_sql_string(res->getString(2))}; + data->set_hard_locality(res->getBoolean(3)); + + std::unique_ptr locality_statement( + static_cast(conn)->prepareStatement( + "SELECT `address` FROM `data_locality` WHERE `id` = ?" + ) + ); + locality_statement->setBytes(1, &id_bytes); + std::unique_ptr const locality_res(locality_statement->executeQuery()); + std::vector locality; + while (locality_res->next()) { + locality.emplace_back(get_sql_string(locality_res->getString(1))); + } + if (!locality.empty()) { + data->set_locality(locality); + } + return StorageErr{}; +} + +auto MySqlDataStorage::get_data(StorageConnection& conn, boost::uuids::uuid const id, Data* data) -> StorageErr { try { - std::unique_ptr statement( + StorageErr const err = get_data_with_locality(conn, id, data); + if (false == err.success()) { + return err; + } + } catch (sql::SQLException& e) { + static_cast(conn)->rollback(); + return StorageErr{StorageErrType::OtherErr, e.what()}; + } + static_cast(conn)->commit(); + return StorageErr{}; +} + +auto MySqlDataStorage::get_driver_data( + StorageConnection& conn, + boost::uuids::uuid const driver_id, + boost::uuids::uuid const data_id, + Data* data +) -> StorageErr { + try { + StorageErr const err = get_data_with_locality(conn, data_id, data); + if (false == err.success()) { + return err; + } + // Add data reference from driver + std::unique_ptr statement{ static_cast(conn)->prepareStatement( - "SELECT `id`, `value`, `hard_locality` FROM `data` WHERE `id` = ?" + "INSERT INTO `data_ref_driver` (`id`, `driver_id`) VALUES (?, ?)" ) - ); - sql::bytes id_bytes = uuid_get_bytes(id); + }; + sql::bytes id_bytes = uuid_get_bytes(data_id); + sql::bytes driver_id_bytes = uuid_get_bytes(driver_id); statement->setBytes(1, &id_bytes); - std::unique_ptr res(statement->executeQuery()); - if (res->rowsCount() == 0) { - static_cast(conn)->rollback(); - return StorageErr{ - StorageErrType::KeyNotFoundErr, - fmt::format("no data with id {}", boost::uuids::to_string(id)) - }; - } - res->next(); - *data = Data{id, get_sql_string(res->getString(2))}; - data->set_hard_locality(res->getBoolean(3)); + statement->setBytes(2, &driver_id_bytes); + statement->executeUpdate(); + } catch (sql::SQLException& e) { + static_cast(conn)->rollback(); + return StorageErr{StorageErrType::OtherErr, e.what()}; + } + static_cast(conn)->commit(); + return StorageErr{}; +} - std::unique_ptr locality_statement( +auto MySqlDataStorage::get_task_data( + StorageConnection& conn, + boost::uuids::uuid const task_id, + boost::uuids::uuid const data_id, + Data* data +) -> StorageErr { + try { + StorageErr const err = get_data_with_locality(conn, data_id, data); + if (false == err.success()) { + return err; + } + // Add data reference from task + std::unique_ptr statement{ static_cast(conn)->prepareStatement( - "SELECT `address` FROM `data_locality` WHERE `id` = ?" + "INSERT INTO `data_ref_task` (`id`, `task_id`) VALUES (?, ?)" ) - ); - locality_statement->setBytes(1, &id_bytes); - std::unique_ptr const locality_res(locality_statement->executeQuery()); - std::vector locality; - while (locality_res->next()) { - locality.emplace_back(get_sql_string(locality_res->getString(1))); - } - if (!locality.empty()) { - data->set_locality(locality); - } + }; + sql::bytes id_bytes = uuid_get_bytes(data_id); + sql::bytes task_id_bytes = uuid_get_bytes(task_id); + statement->setBytes(1, &id_bytes); + statement->setBytes(2, &task_id_bytes); + statement->executeUpdate(); } catch (sql::SQLException& e) { static_cast(conn)->rollback(); return StorageErr{StorageErrType::OtherErr, e.what()}; diff --git a/src/spider/storage/mysql/MySqlStorage.hpp b/src/spider/storage/mysql/MySqlStorage.hpp index 2dcc30f49..6fb2c708d 100644 --- a/src/spider/storage/mysql/MySqlStorage.hpp +++ b/src/spider/storage/mysql/MySqlStorage.hpp @@ -148,6 +148,18 @@ class MySqlDataStorage : public DataStorage { -> StorageErr override; auto get_data(StorageConnection& conn, boost::uuids::uuid id, Data* data) -> StorageErr override; + auto get_driver_data( + StorageConnection& conn, + boost::uuids::uuid driver_id, + boost::uuids::uuid data_id, + Data* data + ) -> StorageErr override; + auto get_task_data( + StorageConnection& conn, + boost::uuids::uuid task_id, + boost::uuids::uuid data_id, + Data* data + ) -> StorageErr override; auto set_data_locality(StorageConnection& conn, Data const& data) -> StorageErr override; auto remove_data(StorageConnection& conn, boost::uuids::uuid id) -> StorageErr override; auto @@ -187,6 +199,9 @@ class MySqlDataStorage : public DataStorage { ) -> StorageErr override; private: + static auto get_data_with_locality(StorageConnection& conn, boost::uuids::uuid id, Data* data) + -> StorageErr; + MySqlDataStorage() = default; friend class MySqlStorageFactory; diff --git a/src/spider/worker/FunctionManager.hpp b/src/spider/worker/FunctionManager.hpp index 0d25d7775..672451c93 100644 --- a/src/spider/worker/FunctionManager.hpp +++ b/src/spider/worker/FunctionManager.hpp @@ -44,7 +44,8 @@ using ArgsBuffer = msgpack::sbuffer; using ResultBuffer = msgpack::sbuffer; -using Function = std::function; +using Function = std::function< + ResultBuffer(TaskContext& context, boost::uuids::uuid task_id, ArgsBuffer const&)>; using FunctionMap = std::vector>; @@ -235,8 +236,12 @@ inline auto create_args_request(std::vector const& args_buffer template class FunctionInvoker { public: - static auto apply(F const& function, TaskContext& context, ArgsBuffer const& args_buffer) - -> ResultBuffer { + static auto apply( + F const& function, + TaskContext& context, + boost::uuids::uuid const task_id, + ArgsBuffer const& args_buffer + ) -> ResultBuffer { // NOLINTBEGIN(cppcoreguidelines-pro-type-union-access,cppcoreguidelines-pro-bounds-pointer-arithmetic) using ArgsTuple = signature::args_t; using ReturnType = signature::ret_t; @@ -301,7 +306,7 @@ class FunctionInvoker { if constexpr (cIsSpecializationV) { boost::uuids::uuid const data_id = arg.as(); std::unique_ptr data = std::make_unique(); - err = data_store->get_data(*conn, data_id, data.get()); + err = data_store->get_task_data(*conn, task_id, data_id, data.get()); if (!err.success()) { return; } @@ -365,13 +370,15 @@ class FunctionManager { if (m_function_map.cend() != get(name)) { return false; } + m_function_map.emplace_back( name, std::bind( &FunctionInvoker::apply, std::move(f), std::placeholders::_1, - std::placeholders::_2 + std::placeholders::_2, + std::placeholders::_3 ) ); return true; diff --git a/src/spider/worker/task_executor.cpp b/src/spider/worker/task_executor.cpp index b5937dc12..f64916399 100644 --- a/src/spider/worker/task_executor.cpp +++ b/src/spider/worker/task_executor.cpp @@ -172,7 +172,7 @@ auto main(int const argc, char** argv) -> int { metadata_store, storage_factory ); - msgpack::sbuffer const result_buffer = (*function)(task_context, args_buffer); + msgpack::sbuffer const result_buffer = (*function)(task_context, task_id, args_buffer); spdlog::debug("Function executed"); // Write result buffer to stdout diff --git a/tests/worker/test-FunctionManager.cpp b/tests/worker/test-FunctionManager.cpp index 7da329306..2f3b18d5a 100644 --- a/tests/worker/test-FunctionManager.cpp +++ b/tests/worker/test-FunctionManager.cpp @@ -15,7 +15,9 @@ #include "../../src/spider/client/TaskContext.hpp" #include "../../src/spider/core/Driver.hpp" #include "../../src/spider/core/Error.hpp" +#include "../../src/spider/core/Task.hpp" #include "../../src/spider/core/TaskContextImpl.hpp" +#include "../../src/spider/core/TaskGraph.hpp" #include "../../src/spider/io/MsgPack.hpp" // IWYU pragma: keep #include "../../src/spider/storage/DataStorage.hpp" #include "../../src/spider/storage/MetadataStorage.hpp" @@ -82,8 +84,9 @@ TEMPLATE_LIST_TEST_CASE( = storage_factory->provide_data_storage(); boost::uuids::random_generator gen; + boost::uuids::uuid const task_id = gen(); spider::TaskContext context = spider::core::TaskContextImpl::create_task_context( - gen(), + task_id, std::move(data_storage), std::move(metadata_storage), std::move(storage_factory) @@ -101,14 +104,14 @@ TEMPLATE_LIST_TEST_CASE( // Run function with two ints should succeed spider::core::ArgsBuffer const args_buffers = spider::core::create_args_buffers(2, 3); constexpr int cExpected = 2 + 3; - msgpack::sbuffer const result = (*function)(context, args_buffers); + msgpack::sbuffer const result = (*function)(context, task_id, args_buffers); msgpack::sbuffer buffer{}; msgpack::pack(buffer, cExpected); REQUIRE(cExpected == spider::core::response_get_result(result).value_or(0)); // Run function with wrong number of inputs should fail spider::core::ArgsBuffer wrong_args_buffers = spider::core::create_args_buffers(1); - msgpack::sbuffer wrong_result = (*function)(context, wrong_args_buffers); + msgpack::sbuffer wrong_result = (*function)(context, task_id, wrong_args_buffers); std::optional> wrong_result_option = spider::core::response_get_error(wrong_result); REQUIRE(wrong_result_option.has_value()); @@ -119,7 +122,7 @@ TEMPLATE_LIST_TEST_CASE( // Run function with wrong type of inputs should fail wrong_args_buffers = spider::core::create_args_buffers(0, "test"); - wrong_result = (*function)(context, wrong_args_buffers); + wrong_result = (*function)(context, task_id, wrong_args_buffers); wrong_result_option = spider::core::response_get_error(wrong_result); REQUIRE(wrong_result_option.has_value()); if (wrong_result_option.has_value()) { @@ -141,8 +144,9 @@ TEMPLATE_LIST_TEST_CASE( = storage_factory->provide_data_storage(); boost::uuids::random_generator gen; + boost::uuids::uuid const task_id = gen(); spider::TaskContext context = spider::core::TaskContextImpl::create_task_context( - gen(), + task_id, std::move(data_storage), std::move(metadata_storage), std::move(storage_factory) @@ -153,7 +157,7 @@ TEMPLATE_LIST_TEST_CASE( spider::core::Function const* function = manager.get_function("tuple_ret_test"); spider::core::ArgsBuffer const args_buffers = spider::core::create_args_buffers("test", 3); - msgpack::sbuffer const result = (*function)(context, args_buffers); + msgpack::sbuffer const result = (*function)(context, task_id, args_buffers); REQUIRE(std::make_tuple("test", 3) == spider::core::response_get_result(result).value_or( std::make_tuple("", 0) @@ -186,8 +190,20 @@ TEMPLATE_LIST_TEST_CASE( REQUIRE(metadata_storage->add_driver(*conn, driver).success()); REQUIRE(data_storage->add_driver_data(*conn, driver_id, data).success()); + boost::uuids::uuid const task_id = gen(); + // Submit a job for valid task id + spider::core::Task task{"data_test"}; + task.set_id(task_id); + task.add_input(spider::core::TaskInput{data.get_id()}); + spider::core::TaskGraph graph; + graph.add_task(task); + graph.add_input_task(task_id); + graph.add_output_task(task_id); + boost::uuids::uuid const job_id = gen(); + REQUIRE(metadata_storage->add_job(*conn, job_id, driver_id, graph).success()); + spider::TaskContext context = spider::core::TaskContextImpl::create_task_context( - gen(), + task_id, data_storage, metadata_storage, storage_factory @@ -198,9 +214,10 @@ TEMPLATE_LIST_TEST_CASE( spider::core::Function const* function = manager.get_function("data_test"); spider::core::ArgsBuffer const args_buffers = spider::core::create_args_buffers(data.get_id()); - msgpack::sbuffer const result = (*function)(context, args_buffers); + msgpack::sbuffer const result = (*function)(context, task_id, args_buffers); REQUIRE(3 == spider::core::response_get_result(result).value_or(0)); + REQUIRE(metadata_storage->remove_job(*conn, job_id).success()); REQUIRE(data_storage->remove_data(*conn, data.get_id()).success()); } } // namespace diff --git a/tests/worker/test-TaskExecutor.cpp b/tests/worker/test-TaskExecutor.cpp index 2df25cf23..ea5b0f6df 100644 --- a/tests/worker/test-TaskExecutor.cpp +++ b/tests/worker/test-TaskExecutor.cpp @@ -19,6 +19,8 @@ #include "../../src/spider/core/Data.hpp" #include "../../src/spider/core/Driver.hpp" #include "../../src/spider/core/Error.hpp" +#include "../../src/spider/core/Task.hpp" +#include "../../src/spider/core/TaskGraph.hpp" #include "../../src/spider/io/BoostAsio.hpp" // IWYU pragma: keep #include "../../src/spider/io/MsgPack.hpp" // IWYU pragma: keep #include "../../src/spider/storage/DataStorage.hpp" @@ -179,6 +181,18 @@ TEMPLATE_LIST_TEST_CASE( REQUIRE(metadata_storage->add_driver(*conn, driver).success()); REQUIRE(data_storage->add_driver_data(*conn, driver_id, data).success()); + // Submit a job for a valid task id + boost::uuids::uuid const task_id = gen(); + spider::core::Task task{"data_test"}; + task.set_id(task_id); + task.add_input(spider::core::TaskInput{data.get_id()}); + spider::core::TaskGraph graph; + graph.add_task(task); + graph.add_input_task(task_id); + graph.add_output_task(task_id); + boost::uuids::uuid const job_id = gen(); + REQUIRE(metadata_storage->add_job(*conn, job_id, driver_id, graph).success()); + absl::flat_hash_map< boost::process::v2::environment::key, boost::process::v2::environment::value> const environment_variable @@ -189,7 +203,7 @@ TEMPLATE_LIST_TEST_CASE( spider::worker::TaskExecutor executor{ context, "data_test", - gen(), + task_id, spider::test::get_storage_url(), get_libraries(), environment_variable, @@ -205,6 +219,7 @@ TEMPLATE_LIST_TEST_CASE( } // Clean up + REQUIRE(metadata_storage->remove_job(*conn, job_id).success()); REQUIRE(data_storage->remove_data(*conn, data.get_id()).success()); }