diff --git a/src/spider/CMakeLists.txt b/src/spider/CMakeLists.txt index 11f39e433..88c74d04f 100644 --- a/src/spider/CMakeLists.txt +++ b/src/spider/CMakeLists.txt @@ -1,6 +1,8 @@ # set variable as CACHE INTERNAL to access it from other scope set(SPIDER_CORE_SOURCES storage/mysql/MySqlConnection.cpp + storage/mysql/MySqlStorageFactory.cpp + storage/mysql/MySqlJobSubmissionBatch.cpp storage/mysql/MySqlStorage.cpp worker/FunctionManager.cpp worker/FunctionNameManager.cpp @@ -27,9 +29,11 @@ set(SPIDER_CORE_HEADERS storage/StorageConnection.hpp storage/mysql/mysql_stmt.hpp storage/mysql/MySqlConnection.hpp + storage/mysql/MySqlStorageFactory.hpp storage/mysql/MySqlStorage.hpp storage/mysql/MySqlJobSubmissionBatch.hpp storage/JobSubmissionBatch.hpp + storage/StorageFactory.hpp worker/FunctionManager.hpp worker/FunctionNameManager.hpp CACHE INTERNAL @@ -91,6 +95,7 @@ target_link_libraries( Boost::program_options Boost::system ${CMAKE_DL_LIBS} + fmt::fmt spdlog::spdlog ) @@ -107,6 +112,7 @@ target_link_libraries( Boost::program_options Boost::system ${CMAKE_DL_LIBS} + fmt::fmt spdlog::spdlog ) add_dependencies(spider_worker spider_task_executor) @@ -132,6 +138,7 @@ target_link_libraries( Boost::headers Boost::program_options absl::flat_hash_map + fmt::fmt spdlog::spdlog ) diff --git a/src/spider/client/Data.hpp b/src/spider/client/Data.hpp index 07cd1c72f..549674f52 100644 --- a/src/spider/client/Data.hpp +++ b/src/spider/client/Data.hpp @@ -15,7 +15,8 @@ #include "../io/MsgPack.hpp" // IWYU pragma: keep #include "../io/Serializer.hpp" #include "../storage/DataStorage.hpp" -#include "../storage/mysql/MySqlConnection.hpp" +#include "../storage/StorageConnection.hpp" +#include "../storage/StorageFactory.hpp" #include "Exception.hpp" namespace spider { @@ -66,13 +67,17 @@ class Data { void set_locality(std::vector const& nodes, bool hard) { m_impl->set_locality(nodes); m_impl->set_hard_locality(hard); - std::variant conn_result - = core::MySqlConnection::create(m_data_store->get_url()); + if (nullptr != m_connection) { + m_data_store->set_data_locality(*m_connection, *m_impl); + return; + } + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - auto& conn = std::get(conn_result); - m_data_store->set_data_locality(conn, *m_impl); + auto conn = std::move(std::get>(conn_result)); + m_data_store->set_data_locality(*conn, *m_impl); } class Builder { @@ -116,28 +121,31 @@ class Data { auto data = std::make_unique(std::string{buffer.data(), buffer.size()}); data->set_locality(m_nodes); data->set_hard_locality(m_hard_locality); - std::variant conn_result - = core::MySqlConnection::create(m_data_store->get_url()); - if (std::holds_alternative(conn_result)) { - throw ConnectionException(std::get(conn_result).description); + std::shared_ptr conn = m_connection; + if (nullptr == conn) { + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); + if (std::holds_alternative(conn_result)) { + throw ConnectionException(std::get(conn_result).description); + } + conn = std::move(std::get>(conn_result)); } - auto& conn = std::get(conn_result); core::StorageErr err; switch (m_data_source) { case DataSource::Driver: - err = m_data_store->add_driver_data(conn, m_source_id, *data); + err = m_data_store->add_driver_data(*conn, m_source_id, *data); if (!err.success()) { throw ConnectionException(err.description); } break; case DataSource::TaskContext: - err = m_data_store->add_task_data(conn, m_source_id, *data); + err = m_data_store->add_task_data(*conn, m_source_id, *data); if (!err.success()) { throw ConnectionException(err.description); } break; } - return Data{std::move(data), m_data_store}; + return Data{std::move(data), m_data_store, m_storage_factory, m_connection}; } private: @@ -148,16 +156,32 @@ class Data { Builder(std::shared_ptr data_store, boost::uuids::uuid const source_id, - DataSource const data_source) + DataSource const data_source, + std::shared_ptr storage_factory) : m_data_store{std::move(data_store)}, m_source_id{source_id}, - m_data_source{data_source} {} + m_data_source{data_source}, + m_storage_factory{std::move(storage_factory)} {} + + Builder(std::shared_ptr data_store, + boost::uuids::uuid const source_id, + DataSource const data_source, + std::shared_ptr storage_factory, + std::shared_ptr connection) + : m_data_store{std::move(data_store)}, + m_source_id{source_id}, + m_data_source{data_source}, + m_storage_factory{std::move(storage_factory)}, + m_connection{std::move(connection)} {} std::vector m_nodes; bool m_hard_locality = false; std::function m_cleanup_func; std::shared_ptr m_data_store; + std::shared_ptr m_storage_factory; + std::shared_ptr m_connection = nullptr; + boost::uuids::uuid m_source_id; DataSource m_data_source; @@ -168,14 +192,28 @@ class Data { Data() = default; private: - Data(std::unique_ptr impl, std::shared_ptr data_store) + Data(std::unique_ptr impl, + std::shared_ptr data_store, + std::shared_ptr storage_factory) + : m_impl{std::move(impl)}, + m_data_store{std::move(data_store)}, + m_storage_factory{std::move(storage_factory)} {} + + Data(std::unique_ptr impl, + std::shared_ptr data_store, + std::shared_ptr storage_factory, + std::shared_ptr connection) : m_impl{std::move(impl)}, - m_data_store{std::move(data_store)} {} + m_data_store{std::move(data_store)}, + m_storage_factory{std::move(storage_factory)}, + m_connection{std::move(connection)} {} [[nodiscard]] auto get_impl() const -> std::unique_ptr const& { return m_impl; } std::unique_ptr m_impl; std::shared_ptr m_data_store; + std::shared_ptr m_storage_factory; + std::shared_ptr m_connection = nullptr; friend class core::DataImpl; friend class core::TaskGraphImpl; diff --git a/src/spider/client/Driver.cpp b/src/spider/client/Driver.cpp index b4d76b781..1ee399eff 100644 --- a/src/spider/client/Driver.cpp +++ b/src/spider/client/Driver.cpp @@ -16,27 +16,26 @@ #include "../core/Error.hpp" #include "../core/KeyValueData.hpp" #include "../io/BoostAsio.hpp" // IWYU pragma: keep -#include "../storage/mysql/MySqlConnection.hpp" -#include "../storage/mysql/MySqlStorage.hpp" +#include "../storage/mysql/MySqlStorageFactory.hpp" +#include "../storage/StorageConnection.hpp" #include "Exception.hpp" namespace spider { -Driver::Driver(std::string const& storage_url) { +Driver::Driver(std::string const& storage_url) + : m_storage_factory{std::make_shared(storage_url)} { boost::uuids::random_generator gen; m_id = gen(); - m_metadata_storage = std::make_shared(storage_url); - m_data_storage = std::make_shared(storage_url); + m_metadata_storage = m_storage_factory->provide_metadata_storage(); + m_data_storage = m_storage_factory->provide_data_storage(); - std::variant conn_result - = core::MySqlConnection::create(storage_url); + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - m_conn = std::make_shared( - std::get(std::move(conn_result)) - ); + m_conn = std::move(std::get>(conn_result)); core::StorageErr const err = m_metadata_storage->add_driver(*m_conn, core::Driver{m_id}); if (!err.success()) { @@ -51,14 +50,14 @@ Driver::Driver(std::string const& storage_url) { m_heartbeat_thread = std::jthread([this](std::stop_token stoken) { while (!stoken.stop_requested()) { std::this_thread::sleep_for(std::chrono::seconds(1)); - std::variant conn_result - = core::MySqlConnection::create(m_metadata_storage->get_url()); + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - auto& conn = std::get(conn_result); + auto conn = std::move(std::get>(conn_result)); - core::StorageErr const err = m_metadata_storage->update_heartbeat(conn, m_id); + core::StorageErr const err = m_metadata_storage->update_heartbeat(*conn, m_id); if (!err.success()) { throw ConnectionException(err.description); } @@ -66,17 +65,18 @@ Driver::Driver(std::string const& storage_url) { }); } -Driver::Driver(std::string const& storage_url, boost::uuids::uuid const id) : m_id{id} { - m_metadata_storage = std::make_shared(storage_url); - m_data_storage = std::make_shared(storage_url); - std::variant conn_result - = core::MySqlConnection::create(storage_url); +Driver::Driver(std::string const& storage_url, boost::uuids::uuid const id) + : m_id{id}, + m_storage_factory{std::make_shared(storage_url)} { + m_metadata_storage = m_storage_factory->provide_metadata_storage(); + m_data_storage = m_storage_factory->provide_data_storage(); + + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - m_conn = std::make_shared( - std::get(std::move(conn_result)) - ); + m_conn = std::move(std::get>(conn_result)); core::StorageErr const err = m_metadata_storage->add_driver(*m_conn, core::Driver{m_id}); if (!err.success()) { @@ -91,14 +91,14 @@ Driver::Driver(std::string const& storage_url, boost::uuids::uuid const id) : m_ m_heartbeat_thread = std::jthread([this](std::stop_token stoken) { while (!stoken.stop_requested()) { std::this_thread::sleep_for(std::chrono::seconds(1)); - std::variant conn_result - = core::MySqlConnection::create(m_metadata_storage->get_url()); + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - auto& conn = std::get(conn_result); + auto conn = std::move(std::get>(conn_result)); - core::StorageErr const err = m_metadata_storage->update_heartbeat(conn, m_id); + core::StorageErr const err = m_metadata_storage->update_heartbeat(*conn, m_id); if (!err.success()) { throw ConnectionException(err.description); } diff --git a/src/spider/client/Driver.hpp b/src/spider/client/Driver.hpp index ff5920a6b..13461a0cd 100644 --- a/src/spider/client/Driver.hpp +++ b/src/spider/client/Driver.hpp @@ -18,9 +18,8 @@ #include "../core/TaskGraphImpl.hpp" #include "../io/Serializer.hpp" #include "../storage/JobSubmissionBatch.hpp" -#include "../storage/mysql/MySqlConnection.hpp" -#include "../storage/mysql/MySqlJobSubmissionBatch.hpp" #include "../storage/StorageConnection.hpp" +#include "../storage/StorageFactory.hpp" #include "../worker/FunctionManager.hpp" #include "../worker/FunctionNameManager.hpp" #include "Data.hpp" @@ -83,7 +82,13 @@ class Driver { template auto get_data_builder() -> Data::Builder { using DataBuilder = typename Data::Builder; - return DataBuilder{m_data_storage, m_id, DataBuilder::DataSource::Driver}; + return DataBuilder{ + m_data_storage, + m_id, + DataBuilder::DataSource::Driver, + m_storage_factory, + m_conn + }; } /** @@ -147,10 +152,7 @@ class Driver { if (nullptr != m_batch) { return; } - m_batch = std::make_shared( - // NOLINTNEXTLINE(cppcoreguidelines-pro-type-static-cast-downcast) - static_cast(*m_conn) - ); + m_batch = m_storage_factory->provide_job_submission_batch(*m_conn); } /** @@ -224,7 +226,13 @@ class Driver { } } - return Job{job_id, m_metadata_storage, m_data_storage, m_conn}; + return Job{ + job_id, + m_metadata_storage, + m_data_storage, + m_storage_factory, + m_conn + }; } /** @@ -268,7 +276,13 @@ class Driver { throw ConnectionException(fmt::format("Failed to start job: {}", err.description)); } - return Job{job_id, m_metadata_storage, m_data_storage, m_conn}; + return Job{ + job_id, + m_metadata_storage, + m_data_storage, + m_storage_factory, + m_conn + }; } /** @@ -293,6 +307,7 @@ class Driver { boost::uuids::uuid m_id; std::shared_ptr m_metadata_storage; std::shared_ptr m_data_storage; + std::shared_ptr m_storage_factory; std::shared_ptr m_conn; std::shared_ptr m_batch{nullptr}; std::jthread m_heartbeat_thread; diff --git a/src/spider/client/Job.hpp b/src/spider/client/Job.hpp index 07be6a6c6..a5c7136a6 100644 --- a/src/spider/client/Job.hpp +++ b/src/spider/client/Job.hpp @@ -21,8 +21,8 @@ #include "../core/JobMetadata.hpp" #include "../io/MsgPack.hpp" // IWYU pragma: keep #include "../storage/MetadataStorage.hpp" -#include "../storage/mysql/MySqlConnection.hpp" #include "../storage/StorageConnection.hpp" +#include "../storage/StorageFactory.hpp" #include "Data.hpp" #include "Exception.hpp" #include "task.hpp" @@ -65,13 +65,13 @@ class Job { */ auto wait_complete() -> void { if (nullptr == m_conn) { - std::variant conn_result - = core::MySqlConnection::create(m_data_storage->get_url()); + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - auto& conn = std::get(conn_result); - wait_complete_conn(conn); + auto conn = std::move(std::get>(conn_result)); + wait_complete_conn(*conn); } else { wait_complete_conn(*m_conn); } @@ -93,14 +93,14 @@ class Job { core::StorageErr err; if (nullptr == m_conn) { - std::variant conn_result - = core::MySqlConnection::create(m_data_storage->get_url()); + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - auto& conn = std::get(conn_result); + auto conn = std::move(std::get>(conn_result)); - err = m_metadata_storage->get_job_status(conn, m_id, &status); + err = m_metadata_storage->get_job_status(*conn, m_id, &status); } else { err = m_metadata_storage->get_job_status(*m_conn, m_id, &status); } @@ -131,13 +131,14 @@ class Job { */ auto get_result() -> ReturnType { if (nullptr == m_conn) { - std::variant conn_result - = core::MySqlConnection::create(m_data_storage->get_url()); + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - auto& conn = std::get(conn_result); - return get_result_conn(conn); + auto conn = std::move(std::get>(conn_result)); + + return get_result_conn(*conn); } return get_result_conn(*m_conn); } @@ -158,18 +159,22 @@ class Job { private: Job(boost::uuids::uuid id, std::shared_ptr metadata_storage, - std::shared_ptr data_storage) + std::shared_ptr data_storage, + std::shared_ptr storage_factory) : m_id{id}, m_metadata_storage{std::move(metadata_storage)}, - m_data_storage{std::move(data_storage)} {} + m_data_storage{std::move(data_storage)}, + m_storage_factory{std::move(storage_factory)} {} Job(boost::uuids::uuid id, std::shared_ptr metadata_storage, std::shared_ptr data_storage, + std::shared_ptr storage_factory, std::shared_ptr conn) : m_id{id}, m_metadata_storage{std::move(metadata_storage)}, m_data_storage{std::move(data_storage)}, + m_storage_factory{std::move(storage_factory)}, m_conn{std::move(conn)} {} auto wait_complete_conn(core::StorageConnection& conn) -> void { @@ -302,7 +307,8 @@ class Job { } return core::DataImpl::create_data( std::make_unique(std::move(data)), - m_data_storage + m_data_storage, + m_storage_factory ); } else { if (output.get_type() != typeid(ReturnType).name()) { @@ -330,6 +336,7 @@ class Job { boost::uuids::uuid m_id; std::shared_ptr m_metadata_storage; std::shared_ptr m_data_storage; + std::shared_ptr m_storage_factory; std::shared_ptr m_conn; friend class Driver; diff --git a/src/spider/client/TaskContext.cpp b/src/spider/client/TaskContext.cpp index f34f4f37c..571d4a697 100644 --- a/src/spider/client/TaskContext.cpp +++ b/src/spider/client/TaskContext.cpp @@ -1,7 +1,9 @@ #include "TaskContext.hpp" +#include #include #include +#include #include #include @@ -9,7 +11,7 @@ #include "../core/Error.hpp" #include "../core/KeyValueData.hpp" -#include "../storage/mysql/MySqlConnection.hpp" +#include "../storage/StorageConnection.hpp" #include "Exception.hpp" namespace spider { @@ -19,15 +21,15 @@ auto TaskContext::get_id() const -> boost::uuids::uuid { } auto TaskContext::kv_store_get(std::string const& key) -> std::optional { - std::variant conn_result - = core::MySqlConnection::create(m_data_store->get_url()); + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - auto& conn = std::get(conn_result); + auto conn = std::move(std::get>(conn_result)); std::string value; - core::StorageErr const err = m_data_store->get_task_kv_data(conn, m_task_id, key, &value); + core::StorageErr const err = m_data_store->get_task_kv_data(*conn, m_task_id, key, &value); if (!err.success()) { if (core::StorageErrType::KeyNotFoundErr == err.type) { return std::nullopt; @@ -38,30 +40,31 @@ auto TaskContext::kv_store_get(std::string const& key) -> std::optional void { - std::variant conn_result - = core::MySqlConnection::create(m_data_store->get_url()); + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - auto& conn = std::get(conn_result); + auto conn = std::move(std::get>(conn_result)); core::KeyValueData const kv_data{key, value, m_task_id}; - core::StorageErr const err = m_data_store->add_task_kv_data(conn, kv_data); + core::StorageErr const err = m_data_store->add_task_kv_data(*conn, kv_data); if (!err.success()) { throw ConnectionException(err.description); } } auto TaskContext::get_jobs() -> std::vector { - std::variant conn_result - = core::MySqlConnection::create(m_metadata_store->get_url()); + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - auto& conn = std::get(conn_result); + auto conn = std::move(std::get>(conn_result)); std::vector job_ids; - core::StorageErr const err = m_metadata_store->get_jobs_by_client_id(conn, m_task_id, &job_ids); + core::StorageErr const err + = m_metadata_store->get_jobs_by_client_id(*conn, m_task_id, &job_ids); if (!err.success()) { throw ConnectionException("Failed to get jobs."); } diff --git a/src/spider/client/TaskContext.hpp b/src/spider/client/TaskContext.hpp index ce80c2bec..0cb0ea40d 100644 --- a/src/spider/client/TaskContext.hpp +++ b/src/spider/client/TaskContext.hpp @@ -19,7 +19,8 @@ #include "../core/TaskGraph.hpp" #include "../core/TaskGraphImpl.hpp" #include "../io/Serializer.hpp" -#include "../storage/mysql/MySqlConnection.hpp" +#include "../storage/StorageConnection.hpp" +#include "../storage/StorageFactory.hpp" #include "Data.hpp" #include "Exception.hpp" #include "Job.hpp" @@ -59,7 +60,12 @@ class TaskContext { template auto get_data_builder() -> Data::Builder { using DataBuilder = typename Data::Builder; - return DataBuilder{m_data_store, m_task_id, DataBuilder::DataSource::TaskContext}; + return DataBuilder{ + m_data_store, + m_task_id, + DataBuilder::DataSource::TaskContext, + m_storage_factory + }; } /** @@ -154,18 +160,19 @@ class TaskContext { graph.add_input_task(new_task.get_id()); graph.add_output_task(new_task.get_id()); - std::variant conn_result - = core::MySqlConnection::create(m_data_store->get_url()); + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - auto& conn = std::get(conn_result); - core::StorageErr err = m_metadata_store->add_job(conn, job_id, m_task_id, graph); + auto conn = std::move(std::get>(conn_result)); + + core::StorageErr err = m_metadata_store->add_job(*conn, job_id, m_task_id, graph); if (!err.success()) { throw ConnectionException(fmt::format("Failed to start job: {}", err.description)); } - return Job{job_id, m_metadata_store, m_data_store}; + return Job{job_id, m_metadata_store, m_data_store, m_storage_factory}; } /** @@ -204,19 +211,20 @@ class TaskContext { boost::uuids::random_generator gen; boost::uuids::uuid const job_id = gen(); - std::variant conn_result - = core::MySqlConnection::create(m_data_store->get_url()); + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - auto& conn = std::get(conn_result); + auto conn = std::move(std::get>(conn_result)); + core::StorageErr const err - = m_metadata_store->add_job(conn, job_id, m_task_id, graph.m_impl->get_graph()); + = m_metadata_store->add_job(*conn, job_id, m_task_id, graph.m_impl->get_graph()); if (!err.success()) { throw ConnectionException(fmt::format("Failed to start job: {}", err.description)); } - return Job{job_id, m_metadata_store, m_data_store}; + return Job{job_id, m_metadata_store, m_data_store, m_storage_factory}; } /** @@ -233,11 +241,13 @@ class TaskContext { TaskContext( boost::uuids::uuid const task_id, std::shared_ptr data_store, - std::shared_ptr metadata_store + std::shared_ptr metadata_store, + std::shared_ptr storage_factory ) : m_task_id{task_id}, m_data_store{std::move(data_store)}, - m_metadata_store{std::move(metadata_store)} {} + m_metadata_store{std::move(metadata_store)}, + m_storage_factory{std::move(storage_factory)} {} auto get_data_store() -> std::shared_ptr { return m_data_store; } @@ -247,6 +257,7 @@ class TaskContext { std::shared_ptr m_data_store; std::shared_ptr m_metadata_store; + std::shared_ptr m_storage_factory; friend class core::TaskContextImpl; }; diff --git a/src/spider/core/DataImpl.hpp b/src/spider/core/DataImpl.hpp index 0b1218879..faa26c558 100644 --- a/src/spider/core/DataImpl.hpp +++ b/src/spider/core/DataImpl.hpp @@ -5,16 +5,20 @@ #include #include "../client/Data.hpp" -#include "../core/Data.hpp" +#include "../storage/StorageFactory.hpp" +#include "Data.hpp" namespace spider::core { class DataImpl { public: template - static auto create_data(std::unique_ptr data, std::shared_ptr data_store) - -> spider::Data { - return spider::Data{std::move(data), data_store}; + static auto create_data( + std::unique_ptr data, + std::shared_ptr data_store, + std::shared_ptr storage_factory + ) -> spider::Data { + return spider::Data{std::move(data), data_store, storage_factory}; } template diff --git a/src/spider/core/TaskContextImpl.hpp b/src/spider/core/TaskContextImpl.hpp index 6d3e0fd3d..e8b569c85 100644 --- a/src/spider/core/TaskContextImpl.hpp +++ b/src/spider/core/TaskContextImpl.hpp @@ -8,6 +8,7 @@ #include "../client/TaskContext.hpp" #include "../storage/DataStorage.hpp" #include "../storage/MetadataStorage.hpp" +#include "../storage/StorageFactory.hpp" namespace spider::core { class TaskContextImpl { @@ -15,9 +16,10 @@ class TaskContextImpl { static auto create_task_context( boost::uuids::uuid const& task_id, std::shared_ptr const& data_storage, - std::shared_ptr const& metadata_storage + std::shared_ptr const& metadata_storage, + std::shared_ptr const& storage_factory ) -> TaskContext { - return TaskContext{task_id, data_storage, metadata_storage}; + return TaskContext{task_id, data_storage, metadata_storage, storage_factory}; } static auto get_data_store(TaskContext const& task_context) -> std::shared_ptr { @@ -28,6 +30,11 @@ class TaskContextImpl { ) -> std::shared_ptr { return task_context.m_metadata_store; } + + static auto get_storage_factory(TaskContext const& task_context + ) -> std::shared_ptr { + return task_context.m_storage_factory; + } }; } // namespace spider::core diff --git a/src/spider/io/msgpack_message.cpp b/src/spider/io/msgpack_message.cpp index 4750eda42..0453f4e0e 100644 --- a/src/spider/io/msgpack_message.cpp +++ b/src/spider/io/msgpack_message.cpp @@ -136,6 +136,7 @@ auto receive_message(boost::asio::ip::tcp::socket& socket) -> std::optional(&body_size_vec[1]), body_size_vec.size() - 1); return buffer; } @@ -155,6 +156,7 @@ auto receive_message(boost::asio::ip::tcp::socket& socket) -> std::optional(&body_vec[1]), body_vec.size() - 1); return buffer; } catch (boost::system::system_error& e) { @@ -221,6 +223,7 @@ auto receive_message_async(std::reference_wrapper co_return std::nullopt; } msgpack::sbuffer buffer; + // NOLINTNEXTLINE(bugprone-bitwise-pointer-cast) buffer.write(std::bit_cast(&body_size_vec[1]), body_size_vec.size() - 1); co_return buffer; } @@ -257,6 +260,7 @@ auto receive_message_async(std::reference_wrapper co_return std::nullopt; } msgpack::sbuffer buffer; + // NOLINTNEXTLINE(bugprone-bitwise-pointer-cast) buffer.write(std::bit_cast(&body_vec[1]), body_vec.size() - 1); co_return buffer; } diff --git a/src/spider/scheduler/FifoPolicy.cpp b/src/spider/scheduler/FifoPolicy.cpp index afd7311ab..cb93fab48 100644 --- a/src/spider/scheduler/FifoPolicy.cpp +++ b/src/spider/scheduler/FifoPolicy.cpp @@ -39,7 +39,7 @@ auto FifoPolicy::task_locality_satisfied(spider::core::Task const& task, std::st if (m_data_cache.contains(data_id)) { data = m_data_cache[data_id]; } else { - if (false == m_data_store->get_data(m_conn, data_id, &data).success()) { + if (false == m_data_store->get_data(*m_conn, data_id, &data).success()) { throw std::runtime_error( fmt::format("Data with id {} not exists.", to_string((data_id))) ); @@ -63,7 +63,7 @@ auto FifoPolicy::task_locality_satisfied(spider::core::Task const& task, std::st FifoPolicy::FifoPolicy( std::shared_ptr const& metadata_store, std::shared_ptr const& data_store, - core::StorageConnection& conn + std::shared_ptr const& conn ) : m_metadata_store{metadata_store}, m_data_store{data_store}, @@ -100,9 +100,9 @@ auto FifoPolicy::schedule_next( auto FifoPolicy::fetch_tasks() -> void { m_data_cache.clear(); - m_metadata_store->get_ready_tasks(m_conn, &m_tasks); + m_metadata_store->get_ready_tasks(*m_conn, &m_tasks); std::vector> instances; - m_metadata_store->get_task_timeout(m_conn, &instances); + m_metadata_store->get_task_timeout(*m_conn, &instances); for (auto const& [instance, task] : instances) { m_tasks.emplace_back(task); } @@ -114,7 +114,7 @@ auto FifoPolicy::fetch_tasks() -> void { auto get_task_job_creation_time = [&](boost::uuids::uuid const task_id) -> std::chrono::system_clock::time_point { boost::uuids::uuid job_id; - if (false == m_metadata_store->get_task_job_id(m_conn, task_id, &job_id).success()) { + if (false == m_metadata_store->get_task_job_id(*m_conn, task_id, &job_id).success()) { throw std::runtime_error(fmt::format("Task with id {} not exists.", to_string(task_id)) ); } @@ -122,7 +122,7 @@ auto FifoPolicy::fetch_tasks() -> void { return job_metadata_map[job_id].get_creation_time(); } core::JobMetadata job_metadata; - if (false == m_metadata_store->get_job_metadata(m_conn, job_id, &job_metadata).success()) { + if (false == m_metadata_store->get_job_metadata(*m_conn, job_id, &job_metadata).success()) { throw std::runtime_error(fmt::format("Job with id {} not exists.", to_string(job_id))); } job_metadata_map[job_id] = job_metadata; diff --git a/src/spider/scheduler/FifoPolicy.hpp b/src/spider/scheduler/FifoPolicy.hpp index aab32c244..0629128de 100644 --- a/src/spider/scheduler/FifoPolicy.hpp +++ b/src/spider/scheduler/FifoPolicy.hpp @@ -22,7 +22,7 @@ class FifoPolicy final : public SchedulerPolicy { FifoPolicy( std::shared_ptr const& metadata_store, std::shared_ptr const& data_store, - core::StorageConnection& conn + std::shared_ptr const& conn ); auto schedule_next(boost::uuids::uuid worker_id, std::string const& worker_addr) @@ -34,8 +34,7 @@ class FifoPolicy final : public SchedulerPolicy { std::shared_ptr m_metadata_store; std::shared_ptr m_data_store; - // NOLINTNEXTLINE(cppcoreguidelines-avoid-const-or-ref-data-members) - core::StorageConnection& m_conn; + std::shared_ptr m_conn; std::vector m_tasks; // NOLINTNEXTLINE(misc-include-cleaner) diff --git a/src/spider/scheduler/SchedulerServer.cpp b/src/spider/scheduler/SchedulerServer.cpp index c59009986..dcc0111ee 100644 --- a/src/spider/scheduler/SchedulerServer.cpp +++ b/src/spider/scheduler/SchedulerServer.cpp @@ -30,14 +30,14 @@ SchedulerServer::SchedulerServer( std::shared_ptr policy, std::shared_ptr metadata_store, std::shared_ptr data_store, - core::StorageConnection& conn, + std::shared_ptr conn, core::StopToken& stop_token ) : m_port{port}, m_policy{std::move(policy)}, m_metadata_store{std::move(metadata_store)}, m_data_store{std::move(data_store)}, - m_conn{conn}, + m_conn{std::move(conn)}, m_stop_token{stop_token} { boost::asio::co_spawn(m_context, receive_message(), boost::asio::detached); std::lock_guard const lock{m_mutex}; @@ -139,7 +139,7 @@ auto SchedulerServer::process_message(boost::asio::ip::tcp::socket socket if (request.has_task_id()) { boost::uuids::uuid job_id; core::StorageErr err - = m_metadata_store->get_task_job_id(m_conn, request.get_task_id(), &job_id); + = m_metadata_store->get_task_job_id(*m_conn, request.get_task_id(), &job_id); // It is possible the job is deleted, so we don't need to reset it if (!err.success()) { spdlog::error( @@ -147,7 +147,7 @@ auto SchedulerServer::process_message(boost::asio::ip::tcp::socket socket boost::uuids::to_string(request.get_task_id()) ); } else { - err = m_metadata_store->reset_job(m_conn, job_id); + err = m_metadata_store->reset_job(*m_conn, job_id); if (!err.success()) { spdlog::error("Cannot reset job {}", boost::uuids::to_string(job_id)); co_return; diff --git a/src/spider/scheduler/SchedulerServer.hpp b/src/spider/scheduler/SchedulerServer.hpp index de3011243..73dae91da 100644 --- a/src/spider/scheduler/SchedulerServer.hpp +++ b/src/spider/scheduler/SchedulerServer.hpp @@ -28,7 +28,7 @@ class SchedulerServer { std::shared_ptr policy, std::shared_ptr metadata_store, std::shared_ptr data_store, - core::StorageConnection& conn, + std::shared_ptr conn, core::StopToken& stop_token ); @@ -46,7 +46,7 @@ class SchedulerServer { std::shared_ptr m_policy; std::shared_ptr m_metadata_store; std::shared_ptr m_data_store; - core::StorageConnection& m_conn; + std::shared_ptr m_conn; boost::asio::io_context m_context; diff --git a/src/spider/scheduler/scheduler.cpp b/src/spider/scheduler/scheduler.cpp index 25135074c..af00be1f2 100644 --- a/src/spider/scheduler/scheduler.cpp +++ b/src/spider/scheduler/scheduler.cpp @@ -6,6 +6,7 @@ #include #include #include +#include #include #include @@ -24,8 +25,9 @@ #include "../io/BoostAsio.hpp" // IWYU pragma: keep #include "../storage/DataStorage.hpp" #include "../storage/MetadataStorage.hpp" -#include "../storage/mysql/MySqlConnection.hpp" -#include "../storage/mysql/MySqlStorage.hpp" +#include "../storage/mysql/MySqlStorageFactory.hpp" +#include "../storage/StorageConnection.hpp" +#include "../storage/StorageFactory.hpp" #include "../utils/StopToken.hpp" #include "FifoPolicy.hpp" #include "SchedulerPolicy.hpp" @@ -70,6 +72,7 @@ auto parse_args(int const argc, char** argv) -> boost::program_options::variable } auto heartbeat_loop( + std::shared_ptr const& storage_factory, std::shared_ptr const& metadata_store, spider::core::Scheduler const& scheduler, spider::core::StopToken& stop_token @@ -78,19 +81,21 @@ auto heartbeat_loop( while (!stop_token.stop_requested()) { std::this_thread::sleep_for(std::chrono::seconds(1)); spdlog::debug("Updating heartbeat"); - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_store->get_url()); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { spdlog::error( - "Failed to connection to storage: {}", + "Failed to connect to storage: {}", std::get(conn_result).description ); fail_count++; continue; } - auto& conn = std::get(conn_result); + auto conn = std::move(std::get>(conn_result + )); + spider::core::StorageErr const err - = metadata_store->update_heartbeat(conn, scheduler.get_id()); + = metadata_store->update_heartbeat(*conn, scheduler.get_id()); if (!err.success()) { spdlog::error("Failed to update scheduler heartbeat: {}", err.description); fail_count++; @@ -105,6 +110,7 @@ auto heartbeat_loop( } auto cleanup_loop( + std::shared_ptr const& storage_factory, std::shared_ptr const& metadata_store, std::shared_ptr const& data_store, spider::core::Scheduler const& scheduler, @@ -113,25 +119,27 @@ auto cleanup_loop( while (!stop_token.stop_requested()) { std::this_thread::sleep_for(std::chrono::seconds(cCleanupInterval)); spdlog::debug("Starting cleanup"); - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_store->get_url()); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { spdlog::error( - "Failed to connection to storage: {}", + "Failed to connect to storage: {}", std::get(conn_result).description ); continue; } - auto& conn = std::get(conn_result); + auto conn = std::move(std::get>(conn_result + )); + spider::core::StorageErr err - = metadata_store->set_scheduler_state(conn, scheduler.get_id(), "gc"); + = metadata_store->set_scheduler_state(*conn, scheduler.get_id(), "gc"); if (!err.success()) { spdlog::error("Failed to set scheduler state to gc: {}", err.description); continue; } - data_store->remove_dangling_data(conn); + data_store->remove_dangling_data(*conn); for (size_t i = 0; i < cRetryCount; ++i) { - err = metadata_store->set_scheduler_state(conn, scheduler.get_id(), "normal"); + err = metadata_store->set_scheduler_state(*conn, scheduler.get_id(), "normal"); if (!err.success()) { spdlog::error("Failed to set scheduler state to normal: {}", err.description); if (i >= cRetryCount - 1) { @@ -184,28 +192,31 @@ auto main(int argc, char** argv) -> int { } // Create storages + std::shared_ptr const storage_factory + = std::make_unique(storage_url); std::shared_ptr const metadata_store - = std::make_shared(storage_url); + = storage_factory->provide_metadata_storage(); std::shared_ptr const data_store - = std::make_shared(storage_url); + = storage_factory->provide_data_storage(); // Initialize storages - std::variant conn_result - = spider::core::MySqlConnection::create(storage_url); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { spdlog::error( "Failed to connection to storage: {}", std::get(conn_result).description ); } - auto& conn = std::get(conn_result); + std::shared_ptr const conn + = std::move(std::get>(conn_result)); - spider::core::StorageErr err = metadata_store->initialize(conn); + spider::core::StorageErr err = metadata_store->initialize(*conn); if (!err.success()) { spdlog::error("Failed to initialize metadata storage: {}", err.description); return cStorageErr; } - err = data_store->initialize(conn); + err = data_store->initialize(*conn); if (!err.success()) { spdlog::error("Failed to initialize data storage: {}", err.description); return cStorageErr; @@ -224,7 +235,7 @@ auto main(int argc, char** argv) -> int { // Register scheduler with storage spider::core::Scheduler const scheduler{scheduler_id, scheduler_addr, port}; - err = metadata_store->add_scheduler(conn, scheduler); + err = metadata_store->add_scheduler(*conn, scheduler); if (!err.success()) { spdlog::error("Failed to register scheduler with storage server: {}", err.description); return cStorageErr; @@ -234,6 +245,7 @@ auto main(int argc, char** argv) -> int { // Start a thread that periodically updates the scheduler's heartbeat std::thread heartbeat_thread{ heartbeat_loop, + std::cref(storage_factory), std::cref(metadata_store), std::ref(scheduler), std::ref(stop_token), @@ -242,6 +254,7 @@ auto main(int argc, char** argv) -> int { // Start a thread that periodically starts cleanup std::thread cleanup_thread{ cleanup_loop, + std::cref(storage_factory), std::cref(metadata_store), std::cref(data_store), std::cref(scheduler), diff --git a/src/spider/storage/DataStorage.hpp b/src/spider/storage/DataStorage.hpp index 610750ae5..68d47b821 100644 --- a/src/spider/storage/DataStorage.hpp +++ b/src/spider/storage/DataStorage.hpp @@ -14,10 +14,10 @@ namespace spider::core { class DataStorage { public: DataStorage() = default; - DataStorage(DataStorage const&) = delete; - DataStorage(DataStorage&&) = delete; - auto operator=(DataStorage const&) -> DataStorage& = delete; - auto operator=(DataStorage&&) -> DataStorage& = delete; + DataStorage(DataStorage const&) = default; + auto operator=(DataStorage const&) -> DataStorage& = default; + DataStorage(DataStorage&&) = default; + auto operator=(DataStorage&&) -> DataStorage& = default; virtual ~DataStorage() = default; virtual auto initialize(StorageConnection& conn) -> StorageErr = 0; @@ -74,8 +74,6 @@ class DataStorage { std::string const& key, std::string* value ) -> StorageErr = 0; - - [[nodiscard]] virtual auto get_url() const -> std::string const& = 0; }; } // namespace spider::core diff --git a/src/spider/storage/JobSubmissionBatch.hpp b/src/spider/storage/JobSubmissionBatch.hpp index 30852a947..06d51bbc8 100644 --- a/src/spider/storage/JobSubmissionBatch.hpp +++ b/src/spider/storage/JobSubmissionBatch.hpp @@ -12,8 +12,8 @@ class JobSubmissionBatch { JobSubmissionBatch() = default; JobSubmissionBatch(JobSubmissionBatch const&) = delete; auto operator=(JobSubmissionBatch const&) -> JobSubmissionBatch& = delete; - JobSubmissionBatch(JobSubmissionBatch&&) = delete; - auto operator=(JobSubmissionBatch&&) -> JobSubmissionBatch& = delete; + JobSubmissionBatch(JobSubmissionBatch&&) = default; + auto operator=(JobSubmissionBatch&&) -> JobSubmissionBatch& = default; virtual ~JobSubmissionBatch() = default; }; } // namespace spider::core diff --git a/src/spider/storage/MetadataStorage.hpp b/src/spider/storage/MetadataStorage.hpp index d1a1d169d..7e7645eb8 100644 --- a/src/spider/storage/MetadataStorage.hpp +++ b/src/spider/storage/MetadataStorage.hpp @@ -19,10 +19,10 @@ namespace spider::core { class MetadataStorage { public: MetadataStorage() = default; - MetadataStorage(MetadataStorage const&) = delete; - MetadataStorage(MetadataStorage&&) = delete; - auto operator=(MetadataStorage const&) -> MetadataStorage& = delete; - auto operator=(MetadataStorage&&) -> MetadataStorage& = delete; + MetadataStorage(MetadataStorage const&) = default; + auto operator=(MetadataStorage const&) -> MetadataStorage& = default; + MetadataStorage(MetadataStorage&&) = default; + auto operator=(MetadataStorage&&) -> MetadataStorage& = default; virtual ~MetadataStorage() = default; virtual auto initialize(StorageConnection& conn) -> StorageErr = 0; @@ -135,8 +135,6 @@ class MetadataStorage { boost::uuids::uuid id, std::string const& state ) -> StorageErr = 0; - - [[nodiscard]] virtual auto get_url() const -> std::string const& = 0; }; } // namespace spider::core diff --git a/src/spider/storage/StorageConnection.hpp b/src/spider/storage/StorageConnection.hpp index 376aadaf1..c1ae7f9d9 100644 --- a/src/spider/storage/StorageConnection.hpp +++ b/src/spider/storage/StorageConnection.hpp @@ -3,7 +3,15 @@ namespace spider::core { -class StorageConnection {}; +class StorageConnection { +public: + StorageConnection() = default; + StorageConnection(StorageConnection const&) = delete; + auto operator=(StorageConnection const&) -> StorageConnection& = delete; + StorageConnection(StorageConnection&&) = default; + auto operator=(StorageConnection&&) -> StorageConnection& = default; + virtual ~StorageConnection() = default; +}; } // namespace spider::core diff --git a/src/spider/storage/StorageFactory.hpp b/src/spider/storage/StorageFactory.hpp new file mode 100644 index 000000000..973c8e19e --- /dev/null +++ b/src/spider/storage/StorageFactory.hpp @@ -0,0 +1,33 @@ +#ifndef SPIDER_STORAGE_STORAGEFACTORY_HPP +#define SPIDER_STORAGE_STORAGEFACTORY_HPP + +#include +#include + +#include "../core/Error.hpp" +#include "DataStorage.hpp" +#include "JobSubmissionBatch.hpp" +#include "MetadataStorage.hpp" +#include "StorageConnection.hpp" + +namespace spider::core { +class StorageFactory { +public: + virtual auto provide_data_storage() -> std::unique_ptr = 0; + virtual auto provide_metadata_storage() -> std::unique_ptr = 0; + virtual auto provide_storage_connection( + ) -> std::variant, StorageErr> = 0; + virtual auto + provide_job_submission_batch(StorageConnection&) -> std::unique_ptr = 0; + + StorageFactory() = default; + StorageFactory(StorageFactory const&) = default; + auto operator=(StorageFactory const&) -> StorageFactory& = default; + StorageFactory(StorageFactory&&) = default; + auto operator=(StorageFactory&&) -> StorageFactory& = default; + virtual ~StorageFactory() = default; +}; + +} // namespace spider::core + +#endif diff --git a/src/spider/storage/mysql/MySqlConnection.cpp b/src/spider/storage/mysql/MySqlConnection.cpp index 6cc56c1f0..7f6def7ff 100644 --- a/src/spider/storage/mysql/MySqlConnection.cpp +++ b/src/spider/storage/mysql/MySqlConnection.cpp @@ -13,10 +13,12 @@ #include #include "../../core/Error.hpp" +#include "../StorageConnection.hpp" namespace spider::core { -auto MySqlConnection::create(std::string const& url) -> std::variant { +auto MySqlConnection::create(std::string const& url +) -> std::variant, StorageErr> { // Validate jdbc url std::regex const url_regex(R"(jdbc:mariadb://[^?]+(\?user=([^&]*)(&password=([^&]*))?)?)"); std::smatch match; @@ -27,7 +29,7 @@ auto MySqlConnection::create(std::string const& url) -> std::variant conn{sql::DriverManager::getConnection(url, properties)}; conn->setAutoCommit(false); - return MySqlConnection{std::move(conn)}; + return std::unique_ptr(new MySqlConnection{std::move(conn)}); } catch (sql::SQLException& e) { return StorageErr{StorageErrType::ConnectionErr, e.what()}; } diff --git a/src/spider/storage/mysql/MySqlConnection.hpp b/src/spider/storage/mysql/MySqlConnection.hpp index 5c9cb2cc7..4c707b9a0 100644 --- a/src/spider/storage/mysql/MySqlConnection.hpp +++ b/src/spider/storage/mysql/MySqlConnection.hpp @@ -13,11 +13,12 @@ namespace spider::core { +// Forward declaration for friend class +class MySqlStorageFactory; + // RAII class for MySQL connection class MySqlConnection : public StorageConnection { public: - static auto create(std::string const& url) -> std::variant; - // Delete copy constructor and copy assignment operator MySqlConnection(MySqlConnection const&) = delete; auto operator=(MySqlConnection const&) -> MySqlConnection& = delete; @@ -25,15 +26,20 @@ class MySqlConnection : public StorageConnection { MySqlConnection(MySqlConnection&&) = default; auto operator=(MySqlConnection&&) -> MySqlConnection& = default; - ~MySqlConnection(); + ~MySqlConnection() override; auto operator*() const -> sql::Connection&; auto operator->() const -> sql::Connection*; private: + static auto create(std::string const& url + ) -> std::variant, StorageErr>; + explicit MySqlConnection(std::unique_ptr conn) : m_connection{std::move(conn)} {}; std::unique_ptr m_connection; + + friend class MySqlStorageFactory; }; } // namespace spider::core diff --git a/src/spider/storage/mysql/MySqlJobSubmissionBatch.cpp b/src/spider/storage/mysql/MySqlJobSubmissionBatch.cpp new file mode 100644 index 000000000..8c54c713f --- /dev/null +++ b/src/spider/storage/mysql/MySqlJobSubmissionBatch.cpp @@ -0,0 +1,63 @@ +#include "MySqlJobSubmissionBatch.hpp" + +#include +#include +#include + +#include "../../core/Error.hpp" +#include "../StorageConnection.hpp" +#include "mysql_stmt.hpp" +#include "MySqlConnection.hpp" + +namespace spider::core { + +// NOLINTBEGIN(cppcoreguidelines-pro-type-static-cast-downcast) +MySqlJobSubmissionBatch::MySqlJobSubmissionBatch(StorageConnection& conn) + : m_job_stmt{static_cast(conn)->prepareStatement(mysql::cInsertJob)}, + m_task_stmt{static_cast(conn)->prepareStatement(mysql::cInsertTask)}, + m_task_input_output_stmt{static_cast(conn)->prepareStatement( + mysql::cInsertTaskInputOutput + )}, + m_task_input_value_stmt{static_cast(conn)->prepareStatement( + mysql::cInsertTaskInputValue + )}, + m_task_input_data_stmt{ + static_cast(conn)->prepareStatement(mysql::cInsertTaskInputData) + }, + m_task_output_stmt{ + static_cast(conn)->prepareStatement(mysql::cInsertTaskOutput) + }, + m_task_dependency_stmt{static_cast(conn)->prepareStatement( + mysql::cInsertTaskDependency + )}, + m_input_task_stmt{ + static_cast(conn)->prepareStatement(mysql::cInsertInputTask) + }, + m_output_task_stmt{ + static_cast(conn)->prepareStatement(mysql::cInsertOutputTask) + } {} + +// NOLINTEND(cppcoreguidelines-pro-type-static-cast-downcast) + +auto MySqlJobSubmissionBatch::submit_batch(StorageConnection& conn) -> StorageErr { + try { + m_job_stmt->executeBatch(); + m_task_stmt->executeBatch(); + m_task_output_stmt->executeBatch(); // Update task outputs in case of input reference + m_task_input_output_stmt->executeBatch(); + m_task_input_value_stmt->executeBatch(); + m_task_input_data_stmt->executeBatch(); + m_task_dependency_stmt->executeBatch(); + m_input_task_stmt->executeBatch(); + m_output_task_stmt->executeBatch(); + } catch (sql::SQLException& e) { + // NOLINTNEXTLINE(cppcoreguidelines-pro-type-static-cast-downcast) + static_cast(conn)->rollback(); + return StorageErr{StorageErrType::OtherErr, e.what()}; + } + // NOLINTNEXTLINE(cppcoreguidelines-pro-type-static-cast-downcast) + static_cast(conn)->commit(); + return StorageErr{}; +} + +} // namespace spider::core diff --git a/src/spider/storage/mysql/MySqlJobSubmissionBatch.hpp b/src/spider/storage/mysql/MySqlJobSubmissionBatch.hpp index 327543fa2..02a7929fb 100644 --- a/src/spider/storage/mysql/MySqlJobSubmissionBatch.hpp +++ b/src/spider/storage/mysql/MySqlJobSubmissionBatch.hpp @@ -3,61 +3,26 @@ #include -#include -#include #include #include "../../core/Error.hpp" #include "../JobSubmissionBatch.hpp" #include "../StorageConnection.hpp" -#include "mysql_stmt.hpp" -#include "MySqlConnection.hpp" namespace spider::core { + +// Forward declaration for friend class +class MySqlStorageFactory; + class MySqlJobSubmissionBatch : public JobSubmissionBatch { public: - explicit MySqlJobSubmissionBatch(sql::Connection& conn) - : m_job_stmt{conn.prepareStatement(mysql::cInsertJob)}, - m_task_stmt{conn.prepareStatement(mysql::cInsertTask)}, - m_task_input_output_stmt{conn.prepareStatement(mysql::cInsertTaskInputOutput)}, - m_task_input_value_stmt{conn.prepareStatement(mysql::cInsertTaskInputValue)}, - m_task_input_data_stmt{conn.prepareStatement(mysql::cInsertTaskInputData)}, - m_task_output_stmt{conn.prepareStatement(mysql::cInsertTaskOutput)}, - m_task_dependency_stmt{conn.prepareStatement(mysql::cInsertTaskDependency)}, - m_input_task_stmt{conn.prepareStatement(mysql::cInsertInputTask)}, - m_output_task_stmt{conn.prepareStatement(mysql::cInsertOutputTask)} {} - - explicit MySqlJobSubmissionBatch(MySqlConnection& conn) - : m_job_stmt{conn->prepareStatement(mysql::cInsertJob)}, - m_task_stmt{conn->prepareStatement(mysql::cInsertTask)}, - m_task_input_output_stmt{conn->prepareStatement(mysql::cInsertTaskInputOutput)}, - m_task_input_value_stmt{conn->prepareStatement(mysql::cInsertTaskInputValue)}, - m_task_input_data_stmt{conn->prepareStatement(mysql::cInsertTaskInputData)}, - m_task_output_stmt{conn->prepareStatement(mysql::cInsertTaskOutput)}, - m_task_dependency_stmt{conn->prepareStatement(mysql::cInsertTaskDependency)}, - m_input_task_stmt{conn->prepareStatement(mysql::cInsertInputTask)}, - m_output_task_stmt{conn->prepareStatement(mysql::cInsertOutputTask)} {} - - auto submit_batch(StorageConnection& conn) -> StorageErr override { - try { - m_job_stmt->executeBatch(); - m_task_stmt->executeBatch(); - m_task_output_stmt->executeBatch(); // Update task outputs in case of input reference - m_task_input_output_stmt->executeBatch(); - m_task_input_value_stmt->executeBatch(); - m_task_input_data_stmt->executeBatch(); - m_task_dependency_stmt->executeBatch(); - m_input_task_stmt->executeBatch(); - m_output_task_stmt->executeBatch(); - } catch (sql::SQLException& e) { - // NOLINTNEXTLINE(cppcoreguidelines-pro-type-static-cast-downcast) - static_cast(conn)->rollback(); - return StorageErr{StorageErrType::OtherErr, e.what()}; - } - // NOLINTNEXTLINE(cppcoreguidelines-pro-type-static-cast-downcast) - static_cast(conn)->commit(); - return StorageErr{}; - } + MySqlJobSubmissionBatch(MySqlJobSubmissionBatch const&) = delete; + auto operator=(MySqlJobSubmissionBatch const&) -> MySqlJobSubmissionBatch& = delete; + MySqlJobSubmissionBatch(MySqlJobSubmissionBatch&&) = default; + auto operator=(MySqlJobSubmissionBatch&&) -> MySqlJobSubmissionBatch& = default; + ~MySqlJobSubmissionBatch() override = default; + + auto submit_batch(StorageConnection& conn) -> StorageErr override; auto get_job_stmt() -> sql::PreparedStatement& { return *m_job_stmt; } @@ -80,6 +45,8 @@ class MySqlJobSubmissionBatch : public JobSubmissionBatch { auto get_output_task_stmt() -> sql::PreparedStatement& { return *m_output_task_stmt; } private: + explicit MySqlJobSubmissionBatch(StorageConnection& conn); + std::unique_ptr m_job_stmt; std::unique_ptr m_task_stmt; std::unique_ptr m_task_input_output_stmt; @@ -89,6 +56,8 @@ class MySqlJobSubmissionBatch : public JobSubmissionBatch { std::unique_ptr m_task_dependency_stmt; std::unique_ptr m_input_task_stmt; std::unique_ptr m_output_task_stmt; + + friend class MySqlStorageFactory; }; } // namespace spider::core diff --git a/src/spider/storage/mysql/MySqlStorage.hpp b/src/spider/storage/mysql/MySqlStorage.hpp index b226e2b96..fc12c3658 100644 --- a/src/spider/storage/mysql/MySqlStorage.hpp +++ b/src/spider/storage/mysql/MySqlStorage.hpp @@ -5,7 +5,6 @@ #include #include #include -#include #include #include @@ -27,16 +26,16 @@ #include "MySqlJobSubmissionBatch.hpp" namespace spider::core { -class MySqlMetadataStorage : public MetadataStorage { -public: - MySqlMetadataStorage() = delete; - explicit MySqlMetadataStorage(std::string url) : m_url{std::move(url)} {} +// Forward declaration for friend class +class MySqlStorageFactory; - MySqlMetadataStorage(MySqlMetadataStorage const&) = delete; - MySqlMetadataStorage(MySqlMetadataStorage&&) = delete; - auto operator=(MySqlMetadataStorage const&) -> MySqlMetadataStorage& = delete; - auto operator=(MySqlMetadataStorage&&) -> MySqlMetadataStorage& = delete; +class MySqlMetadataStorage : public MetadataStorage { +public: + MySqlMetadataStorage(MySqlMetadataStorage const&) = default; + MySqlMetadataStorage(MySqlMetadataStorage&&) = default; + auto operator=(MySqlMetadataStorage const&) -> MySqlMetadataStorage& = default; + auto operator=(MySqlMetadataStorage&&) -> MySqlMetadataStorage& = default; ~MySqlMetadataStorage() override = default; auto initialize(StorageConnection& conn) -> StorageErr override; auto add_driver(StorageConnection& conn, Driver const& driver) -> StorageErr override; @@ -128,10 +127,8 @@ class MySqlMetadataStorage : public MetadataStorage { std::string const& state ) -> StorageErr override; - [[nodiscard]] auto get_url() const -> std::string const& override { return m_url; } - private: - std::string m_url; + MySqlMetadataStorage() = default; static void add_task( MySqlConnection& conn, @@ -147,18 +144,16 @@ class MySqlMetadataStorage : public MetadataStorage { ); static auto fetch_full_task(MySqlConnection& conn, std::unique_ptr const& res) -> Task; + + friend class MySqlStorageFactory; }; class MySqlDataStorage : public DataStorage { public: - MySqlDataStorage() = delete; - - explicit MySqlDataStorage(std::string url) : m_url{std::move(url)} {} - - MySqlDataStorage(MySqlDataStorage const&) = delete; - MySqlDataStorage(MySqlDataStorage&&) = delete; - auto operator=(MySqlDataStorage const&) -> MySqlDataStorage& = delete; - auto operator=(MySqlDataStorage&&) -> MySqlDataStorage& = delete; + MySqlDataStorage(MySqlDataStorage const&) = default; + MySqlDataStorage(MySqlDataStorage&&) = default; + auto operator=(MySqlDataStorage const&) -> MySqlDataStorage& = default; + auto operator=(MySqlDataStorage&&) -> MySqlDataStorage& = default; ~MySqlDataStorage() override = default; auto initialize(StorageConnection& conn) -> StorageErr override; auto add_driver_data(StorageConnection& conn, boost::uuids::uuid driver_id, Data const& data) @@ -207,10 +202,10 @@ class MySqlDataStorage : public DataStorage { std::string* value ) -> StorageErr override; - [[nodiscard]] auto get_url() const -> std::string const& override { return m_url; } - private: - std::string m_url; + MySqlDataStorage() = default; + + friend class MySqlStorageFactory; }; } // namespace spider::core diff --git a/src/spider/storage/mysql/MySqlStorageFactory.cpp b/src/spider/storage/mysql/MySqlStorageFactory.cpp new file mode 100644 index 000000000..598eaf5f3 --- /dev/null +++ b/src/spider/storage/mysql/MySqlStorageFactory.cpp @@ -0,0 +1,43 @@ +#include "MySqlStorageFactory.hpp" + +#include +#include +#include +#include + +#include "../../core/Error.hpp" +#include "../DataStorage.hpp" +#include "../JobSubmissionBatch.hpp" +#include "../MetadataStorage.hpp" +#include "../StorageConnection.hpp" +#include "MySqlConnection.hpp" +#include "MySqlJobSubmissionBatch.hpp" +#include "MySqlStorage.hpp" + +namespace spider::core { + +MySqlStorageFactory::MySqlStorageFactory(std::string url) : m_url{std::move(url)} {} + +auto MySqlStorageFactory::provide_data_storage() -> std::unique_ptr { + return std::unique_ptr(new MySqlDataStorage()); +} + +auto MySqlStorageFactory::provide_metadata_storage() -> std::unique_ptr { + return std::unique_ptr(new MySqlMetadataStorage()); +} + +auto MySqlStorageFactory::provide_storage_connection( +) -> std::variant, StorageErr> { + std::variant, StorageErr> connection + = MySqlConnection::create(m_url); + if (std::holds_alternative(connection)) { + return std::get(connection); + } + return std::move(std::get>(connection)); +} + +auto MySqlStorageFactory::provide_job_submission_batch(StorageConnection& connection +) -> std::unique_ptr { + return std::unique_ptr(new MySqlJobSubmissionBatch(connection)); +} +} // namespace spider::core diff --git a/src/spider/storage/mysql/MySqlStorageFactory.hpp b/src/spider/storage/mysql/MySqlStorageFactory.hpp new file mode 100644 index 000000000..daf02a8ed --- /dev/null +++ b/src/spider/storage/mysql/MySqlStorageFactory.hpp @@ -0,0 +1,32 @@ +#ifndef SPIDER_STORAGE_MYSQLSTORAGEFACTORY_HPP +#define SPIDER_STORAGE_MYSQLSTORAGEFACTORY_HPP + +#include +#include +#include + +#include "../../core/Error.hpp" +#include "../DataStorage.hpp" +#include "../JobSubmissionBatch.hpp" +#include "../MetadataStorage.hpp" +#include "../StorageConnection.hpp" +#include "../StorageFactory.hpp" + +namespace spider::core { +class MySqlStorageFactory : public StorageFactory { +public: + explicit MySqlStorageFactory(std::string url); + + auto provide_data_storage() -> std::unique_ptr override; + auto provide_metadata_storage() -> std::unique_ptr override; + auto provide_storage_connection( + ) -> std::variant, StorageErr> override; + auto provide_job_submission_batch(StorageConnection&) + -> std::unique_ptr override; + +private: + std::string m_url; +}; +} // namespace spider::core + +#endif diff --git a/src/spider/worker/FunctionManager.hpp b/src/spider/worker/FunctionManager.hpp index 9a226699e..92088cb33 100644 --- a/src/spider/worker/FunctionManager.hpp +++ b/src/spider/worker/FunctionManager.hpp @@ -26,7 +26,7 @@ #include "../io/MsgPack.hpp" // IWYU pragma: keep #include "../io/Serializer.hpp" #include "../storage/DataStorage.hpp" -#include "../storage/mysql/MySqlConnection.hpp" +#include "../storage/StorageConnection.hpp" #include "TaskExecutorMessage.hpp" // NOLINTBEGIN(cppcoreguidelines-macro-usage) @@ -283,8 +283,8 @@ class FunctionInvoker { // Fill args_tuple StorageErr err; std::get<0>(args_tuple) = context; - std::variant conn_result - = core::MySqlConnection::create(data_store->get_url()); + std::variant, core::StorageErr> conn_result + = TaskContextImpl::get_storage_factory(context)->provide_storage_connection(); if (std::holds_alternative(conn_result)) { err = std::get(conn_result); return create_error_response( @@ -292,7 +292,7 @@ class FunctionInvoker { fmt::format("Cannot parse arguments: {}.", err.description) ); } - auto& conn = std::get(conn_result); + auto conn = std::move(std::get>(conn_result)); for_n - 1>([&](auto i) { if (!err.success()) { return; @@ -302,13 +302,17 @@ 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_data(*conn, data_id, data.get()); if (!err.success()) { return; } - std::get(args_tuple - ) = DataImpl::create_data>(std::move(data), data_store); + std::get(args_tuple) + = DataImpl::create_data>( + std::move(data), + data_store, + TaskContextImpl::get_storage_factory(context) + ); } else { std::get(args_tuple) = arg.as>(); diff --git a/src/spider/worker/WorkerClient.cpp b/src/spider/worker/WorkerClient.cpp index b504bdc68..705399ff8 100644 --- a/src/spider/worker/WorkerClient.cpp +++ b/src/spider/worker/WorkerClient.cpp @@ -24,7 +24,8 @@ #include "../scheduler/SchedulerMessage.hpp" #include "../storage/DataStorage.hpp" #include "../storage/MetadataStorage.hpp" -#include "../storage/mysql/MySqlConnection.hpp" +#include "../storage/StorageConnection.hpp" +#include "../storage/StorageFactory.hpp" namespace spider::worker { @@ -32,12 +33,14 @@ WorkerClient::WorkerClient( boost::uuids::uuid const worker_id, std::string worker_addr, std::shared_ptr data_store, - std::shared_ptr metadata_store + std::shared_ptr metadata_store, + std::shared_ptr storage_factory ) : m_worker_id{worker_id}, m_worker_addr{std::move(worker_addr)}, m_data_store(std::move(data_store)), - m_metadata_store(std::move(metadata_store)) {} + m_metadata_store(std::move(metadata_store)), + m_storage_factory(std::move(storage_factory)) {} auto WorkerClient::get_next_task(std::optional const& fail_task_id ) -> std::optional> { @@ -45,8 +48,8 @@ auto WorkerClient::get_next_task(std::optional const& fail_t std::vector schedulers; { // Keep the scope for RAII storage connection - std::variant conn_result - = core::MySqlConnection::create(m_metadata_store->get_url()); + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { spdlog::error( "Failed to connect to storage: {}", @@ -54,8 +57,8 @@ auto WorkerClient::get_next_task(std::optional const& fail_t ); return std::nullopt; } - auto& conn = std::get(conn_result); - if (!m_metadata_store->get_active_scheduler(conn, &schedulers).success()) { + auto conn = std::move(std::get>(conn_result)); + if (!m_metadata_store->get_active_scheduler(*conn, &schedulers).success()) { return std::nullopt; } } @@ -114,8 +117,9 @@ auto WorkerClient::get_next_task(std::optional const& fail_t return std::nullopt; } boost::uuids::uuid const task_id = response.get_task_id(); - std::variant conn_result - = core::MySqlConnection::create(m_metadata_store->get_url()); + + std::variant, core::StorageErr> conn_result + = m_storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { spdlog::error( "Failed to connect to storage: {}", @@ -123,9 +127,10 @@ auto WorkerClient::get_next_task(std::optional const& fail_t ); return std::nullopt; } - auto& conn = std::get(conn_result); + auto conn = std::move(std::get>(conn_result)); + core::TaskInstance const instance{task_id}; - core::StorageErr const err = m_metadata_store->create_task_instance(conn, instance); + core::StorageErr const err = m_metadata_store->create_task_instance(*conn, instance); if (!err.success()) { return std::nullopt; } diff --git a/src/spider/worker/WorkerClient.hpp b/src/spider/worker/WorkerClient.hpp index 523869a7b..7f408c2c4 100644 --- a/src/spider/worker/WorkerClient.hpp +++ b/src/spider/worker/WorkerClient.hpp @@ -11,6 +11,7 @@ #include "../io/BoostAsio.hpp" // IWYU pragma: keep #include "../storage/DataStorage.hpp" #include "../storage/MetadataStorage.hpp" +#include "../storage/StorageFactory.hpp" namespace spider::worker { class WorkerClient { @@ -26,7 +27,8 @@ class WorkerClient { boost::uuids::uuid worker_id, std::string worker_addr, std::shared_ptr data_store, - std::shared_ptr metadata_store + std::shared_ptr metadata_store, + std::shared_ptr storage_factory ); auto get_next_task(std::optional const& fail_task_id @@ -38,6 +40,7 @@ class WorkerClient { std::shared_ptr m_data_store; std::shared_ptr m_metadata_store; + std::shared_ptr m_storage_factory; }; } // namespace spider::worker #endif // SPIDER_WORKER_WORKERCLIENT_HPP diff --git a/src/spider/worker/task_executor.cpp b/src/spider/worker/task_executor.cpp index 3ff5deab9..f92a7fd6b 100644 --- a/src/spider/worker/task_executor.cpp +++ b/src/spider/worker/task_executor.cpp @@ -24,7 +24,8 @@ #include "../io/MsgPack.hpp" // IWYU pragma: keep #include "../storage/DataStorage.hpp" #include "../storage/MetadataStorage.hpp" -#include "../storage/mysql/MySqlStorage.hpp" +#include "../storage/mysql/MySqlStorageFactory.hpp" +#include "../storage/StorageFactory.hpp" #include "DllLoader.hpp" #include "FunctionManager.hpp" #include "message_pipe.hpp" @@ -122,10 +123,12 @@ auto main(int const argc, char** argv) -> int { boost::uuids::uuid const task_id = gen(task_id_string); // Set up storage + std::shared_ptr const storage_factory + = std::make_shared(storage_url); std::shared_ptr const metadata_store - = std::make_shared(storage_url); + = storage_factory->provide_metadata_storage(); std::shared_ptr const data_store - = std::make_shared(storage_url); + = storage_factory->provide_data_storage(); // Set up asio boost::asio::io_context context; @@ -167,7 +170,8 @@ auto main(int const argc, char** argv) -> int { spider::TaskContext task_context = spider::core::TaskContextImpl::create_task_context( task_id, data_store, - metadata_store + metadata_store, + storage_factory ); msgpack::sbuffer const result_buffer = (*function)(task_context, args_buffer); spdlog::debug("Function executed"); diff --git a/src/spider/worker/worker.cpp b/src/spider/worker/worker.cpp index e855746a0..983c28241 100644 --- a/src/spider/worker/worker.cpp +++ b/src/spider/worker/worker.cpp @@ -9,6 +9,7 @@ #include #include #include +#include #include #include @@ -38,8 +39,9 @@ #include "../io/Serializer.hpp" // IWYU pragma: keep #include "../storage/DataStorage.hpp" #include "../storage/MetadataStorage.hpp" -#include "../storage/mysql/MySqlConnection.hpp" -#include "../storage/mysql/MySqlStorage.hpp" +#include "../storage/mysql/MySqlStorageFactory.hpp" +#include "../storage/StorageConnection.hpp" +#include "../storage/StorageFactory.hpp" #include "../utils/StopToken.hpp" #include "TaskExecutor.hpp" #include "WorkerClient.hpp" @@ -100,6 +102,7 @@ auto get_environment_variable() -> absl::flat_hash_map< } auto heartbeat_loop( + std::shared_ptr const& storage_factory, std::shared_ptr const& metadata_store, spider::core::Driver const& driver, spider::core::StopToken& stop_token @@ -108,20 +111,21 @@ auto heartbeat_loop( while (!stop_token.stop_requested()) { std::this_thread::sleep_for(std::chrono::seconds(1)); spdlog::debug("Updating heartbeat"); - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_store->get_url()); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { spdlog::error( - "Failed to connection to storage: {}", + "Failed to connect to storage: {}", std::get(conn_result).description ); fail_count++; continue; } - auto& conn = std::get(conn_result); + auto conn = std::move(std::get>(conn_result + )); spider::core::StorageErr const err - = metadata_store->update_heartbeat(conn, driver.get_id()); + = metadata_store->update_heartbeat(*conn, driver.get_id()); if (!err.success()) { spdlog::error("Failed to update scheduler heartbeat: {}", err.description); fail_count++; @@ -219,8 +223,10 @@ auto parse_outputs( // NOLINTBEGIN(clang-analyzer-unix.BlockInCriticalSection) auto task_loop( + std::shared_ptr const& storage_factory, std::shared_ptr const& metadata_store, spider::worker::WorkerClient& client, + std::string const& storage_url, std::vector const& libs, absl::flat_hash_map< boost::process::v2::environment::key, @@ -242,17 +248,20 @@ auto task_loop( { // Keep the scope of RAII storage connection - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_store->get_url()); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { spdlog::error( - "Failed to connection to storage: {}", + "Failed to connect to storage: {}", std::get(conn_result).description ); continue; } - auto& conn = std::get(conn_result); - err = metadata_store->get_task(conn, task_id, &task); + auto conn = std::move( + std::get>(conn_result) + ); + + err = metadata_store->get_task(*conn, task_id, &task); if (!err.success()) { spdlog::error("Failed to fetch task detail: {}", err.description); continue; @@ -262,7 +271,7 @@ auto task_loop( optional_args_buffers = get_args_buffers(task); if (!optional_args_buffers.has_value()) { metadata_store->task_fail( - conn, + *conn, instance, fmt::format("Task {} failed to parse arguments", task.get_function_name()) ); @@ -276,7 +285,7 @@ auto task_loop( context, task.get_function_name(), task.get_id(), - metadata_store->get_url(), + storage_url, libs, environment, args_buffers @@ -285,20 +294,22 @@ auto task_loop( context.run(); executor.wait(); - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_store->get_url()); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { spdlog::error( - "Failed to connection to storage: {}", + "Failed to connect to storage: {}", std::get(conn_result).description ); continue; } - auto& conn = std::get(conn_result); + auto conn = std::move(std::get>(conn_result + )); + if (!executor.succeed()) { spdlog::warn("Task {} failed", task.get_function_name()); metadata_store->task_fail( - conn, + *conn, instance, fmt::format("Task {} failed", task.get_function_name()) ); @@ -312,7 +323,7 @@ auto task_loop( if (!optional_result_buffers.has_value()) { spdlog::error("Task {} failed to parse result into buffers", task.get_function_name()); metadata_store->task_fail( - conn, + *conn, instance, fmt::format( "Task {} failed to parse result into buffers", @@ -327,7 +338,7 @@ auto task_loop( = parse_outputs(task, result_buffers); if (!optional_outputs.has_value()) { metadata_store->task_fail( - conn, + *conn, instance, fmt::format( "Task {} failed to parse result into TaskOutput", @@ -341,7 +352,7 @@ auto task_loop( // Submit result spdlog::debug("Submitting result for task {}", boost::uuids::to_string(task_id)); for (int i = 0; i < cRetryCount; ++i) { - err = metadata_store->task_finish(conn, instance, outputs); + err = metadata_store->task_finish(*conn, instance, outputs); if (err.success()) { break; } @@ -404,27 +415,31 @@ auto main(int argc, char** argv) -> int { } // Create storage + std::shared_ptr const storage_factory + = std::make_shared(storage_url); std::shared_ptr const metadata_store - = std::make_shared(storage_url); + = storage_factory->provide_metadata_storage(); std::shared_ptr const data_store - = std::make_shared(storage_url); + = storage_factory->provide_data_storage(); boost::uuids::random_generator gen; boost::uuids::uuid const worker_id = gen(); spider::core::Driver driver{worker_id}; { // Keep the scope of RAII storage connection - std::variant conn_result - = spider::core::MySqlConnection::create(storage_url); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); if (std::holds_alternative(conn_result)) { spdlog::error( - "Failed to connection to storage: {}", + "Failed to connect to storage: {}", std::get(conn_result).description ); return cStorageErr; } - auto& conn = std::get(conn_result); - spider::core::StorageErr const err = metadata_store->add_driver(conn, driver); + auto conn = std::move(std::get>(conn_result + )); + + spider::core::StorageErr const err = metadata_store->add_driver(*conn, driver); if (!err.success()) { spdlog::error("Cannot add driver to metadata storage: {}", err.description); return cStorageErr; @@ -434,7 +449,8 @@ auto main(int argc, char** argv) -> int { spider::core::StopToken stop_token; // Start client - spider::worker::WorkerClient client{worker_id, worker_addr, data_store, metadata_store}; + spider::worker::WorkerClient + client{worker_id, worker_addr, data_store, metadata_store, storage_factory}; absl::flat_hash_map< boost::process::v2::environment::key, @@ -444,6 +460,7 @@ auto main(int argc, char** argv) -> int { // Start a thread that periodically updates the scheduler's heartbeat std::thread heartbeat_thread{ heartbeat_loop, + std::cref(storage_factory), std::cref(metadata_store), std::ref(driver), std::ref(stop_token) @@ -452,8 +469,10 @@ auto main(int argc, char** argv) -> int { // Start a thread that processes tasks std::thread task_thread{ task_loop, + std::cref(storage_factory), std::cref(metadata_store), std::ref(client), + std::cref(storage_url), std::cref(libs), std::cref(environment_variables), std::cref(stop_token), diff --git a/tests/client/test-Driver.cpp b/tests/client/test-Driver.cpp index 64d87c170..dd245a35f 100644 --- a/tests/client/test-Driver.cpp +++ b/tests/client/test-Driver.cpp @@ -3,6 +3,7 @@ #include #include +#include #include #include "../../src/spider/client/Data.hpp" @@ -12,8 +13,13 @@ #include "../storage/StorageTestHelper.hpp" namespace { -TEST_CASE("Driver kv store", "[client][storage]") { - spider::Driver driver{spider::test::cStorageUrl}; +TEMPLATE_LIST_TEST_CASE( + "Driver kv store", + "[client][storage]", + spider::test::StorageFactoryTypeList +) { + std::string const storage_url = spider::test::get_storage_url(); + spider::Driver driver{storage_url}; driver.kv_store_insert("key", "value"); // Get value by key should succeed @@ -27,8 +33,9 @@ TEST_CASE("Driver kv store", "[client][storage]") { REQUIRE(!fail_result.has_value()); } -TEST_CASE("Driver data", "[client][storage]") { - spider::Driver driver{spider::test::cStorageUrl}; +TEMPLATE_LIST_TEST_CASE("Driver data", "[client][storage]", spider::test::StorageFactoryTypeList) { + std::string const storage_url = spider::test::get_storage_url(); + spider::Driver driver{storage_url}; spider::Data const data = driver.get_data_builder().build(1); } @@ -43,16 +50,26 @@ auto test_driver(spider::TaskContext&, spider::Data& x) -> int { SPIDER_REGISTER_TASK(sum); SPIDER_REGISTER_TASK(test_driver); -TEST_CASE("Driver bind task", "[client][storage]") { - spider::Driver driver{spider::test::cStorageUrl}; +TEMPLATE_LIST_TEST_CASE( + "Driver bind task", + "[client][storage]", + spider::test::StorageFactoryTypeList +) { + std::string const storage_url = spider::test::get_storage_url(); + spider::Driver driver{storage_url}; spider::TaskGraph const graph_1 = driver.bind(&sum, &sum, 0); spider::TaskGraph const graph_3 = driver.bind(&sum, &sum, &sum); spider::TaskGraph const graph_4 = driver.bind(&sum, graph_1, graph_1); } -TEST_CASE("Driver bind task with data", "[client][storage]") { - spider::Driver driver{spider::test::cStorageUrl}; +TEMPLATE_LIST_TEST_CASE( + "Driver bind task with data", + "[client][storage]", + spider::test::StorageFactoryTypeList +) { + std::string const storage_url = spider::test::get_storage_url(); + spider::Driver driver{storage_url}; spider::Data data = driver.get_data_builder().build(1); spider::TaskGraph const graph_1 = driver.bind(&test_driver, data); diff --git a/tests/scheduler/test-SchedulerPolicy.cpp b/tests/scheduler/test-SchedulerPolicy.cpp index eca670f5c..db4dfb8bd 100644 --- a/tests/scheduler/test-SchedulerPolicy.cpp +++ b/tests/scheduler/test-SchedulerPolicy.cpp @@ -4,7 +4,6 @@ #include #include #include -#include #include #include @@ -21,29 +20,28 @@ #include "../../src/spider/scheduler/FifoPolicy.hpp" #include "../../src/spider/storage/DataStorage.hpp" #include "../../src/spider/storage/MetadataStorage.hpp" -#include "../../src/spider/storage/mysql/MySqlConnection.hpp" +#include "../../src/spider/storage/StorageConnection.hpp" +#include "../../src/spider/storage/StorageFactory.hpp" #include "../storage/StorageTestHelper.hpp" namespace { TEMPLATE_LIST_TEST_CASE( "FIFO schedule order", "[scheduler][storage]", - spider::test::StorageTypeList + spider::test::StorageFactoryTypeList ) { - std::tuple< - std::unique_ptr, - std::unique_ptr> - storages = spider::test::create_storage< - std::tuple_element_t<0, TestType>, - std::tuple_element_t<1, TestType>>(); + std::shared_ptr const storage_factory + = spider::test::create_storage_factory(); std::shared_ptr const metadata_store - = std::move(std::get<0>(storages)); - std::shared_ptr const data_store = std::move(std::get<1>(storages)); + = storage_factory->provide_metadata_storage(); + std::shared_ptr const data_store + = storage_factory->provide_data_storage(); - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_store->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + std::shared_ptr const conn + = std::move(std::get>(conn_result)); boost::uuids::random_generator gen; boost::uuids::uuid const client_id = gen(); @@ -54,7 +52,7 @@ TEMPLATE_LIST_TEST_CASE( graph_1.add_input_task(task_1.get_id()); graph_1.add_output_task(task_1.get_id()); boost::uuids::uuid const job_id_1 = gen(); - REQUIRE(metadata_store->add_job(conn, job_id_1, client_id, graph_1).success()); + REQUIRE(metadata_store->add_job(*conn, job_id_1, client_id, graph_1).success()); std::this_thread::sleep_for(std::chrono::seconds(1)); spider::core::Task const task_2{"task_2"}; spider::core::TaskGraph graph_2; @@ -62,7 +60,7 @@ TEMPLATE_LIST_TEST_CASE( graph_2.add_input_task(task_2.get_id()); graph_2.add_output_task(task_2.get_id()); boost::uuids::uuid const job_id_2 = gen(); - REQUIRE(metadata_store->add_job(conn, job_id_2, client_id, graph_2).success()); + REQUIRE(metadata_store->add_job(*conn, job_id_2, client_id, graph_2).success()); spider::scheduler::FifoPolicy policy{metadata_store, data_store, conn}; @@ -82,8 +80,8 @@ TEMPLATE_LIST_TEST_CASE( REQUIRE(task_id == task_2.get_id()); } - REQUIRE(metadata_store->remove_job(conn, job_id_1).success()); - REQUIRE(metadata_store->remove_job(conn, job_id_2).success()); + REQUIRE(metadata_store->remove_job(*conn, job_id_1).success()); + REQUIRE(metadata_store->remove_job(*conn, job_id_2).success()); // Schedule when no task available optional_task_id = policy.schedule_next(gen(), ""); @@ -93,22 +91,20 @@ TEMPLATE_LIST_TEST_CASE( TEMPLATE_LIST_TEST_CASE( "Schedule hard locality", "[scheduler][storage]", - spider::test::StorageTypeList + spider::test::StorageFactoryTypeList ) { - std::tuple< - std::unique_ptr, - std::unique_ptr> - storages = spider::test::create_storage< - std::tuple_element_t<0, TestType>, - std::tuple_element_t<1, TestType>>(); + std::shared_ptr const storage_factory + = spider::test::create_storage_factory(); std::shared_ptr const metadata_store - = std::move(std::get<0>(storages)); - std::shared_ptr const data_store = std::move(std::get<1>(storages)); + = storage_factory->provide_metadata_storage(); + std::shared_ptr const data_store + = storage_factory->provide_data_storage(); - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_store->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + std::shared_ptr const conn + = std::move(std::get>(conn_result)); boost::uuids::random_generator gen; boost::uuids::uuid const job_id = gen(); @@ -118,14 +114,14 @@ TEMPLATE_LIST_TEST_CASE( spider::core::Data data{"value"}; data.set_hard_locality(true); data.set_locality({"127.0.0.1"}); - REQUIRE(metadata_store->add_driver(conn, spider::core::Driver{client_id}).success()); - REQUIRE(data_store->add_driver_data(conn, client_id, data).success()); + REQUIRE(metadata_store->add_driver(*conn, spider::core::Driver{client_id}).success()); + REQUIRE(data_store->add_driver_data(*conn, client_id, data).success()); task.add_input(spider::core::TaskInput{data.get_id()}); spider::core::TaskGraph graph; graph.add_task(task); graph.add_input_task(task.get_id()); graph.add_output_task(task.get_id()); - REQUIRE(metadata_store->add_job(conn, job_id, client_id, graph).success()); + REQUIRE(metadata_store->add_job(*conn, job_id, client_id, graph).success()); spider::scheduler::FifoPolicy policy{metadata_store, data_store, conn}; // Schedule with wrong address @@ -139,28 +135,26 @@ TEMPLATE_LIST_TEST_CASE( REQUIRE(task_id == task.get_id()); } - REQUIRE(metadata_store->remove_job(conn, job_id).success()); + REQUIRE(metadata_store->remove_job(*conn, job_id).success()); } TEMPLATE_LIST_TEST_CASE( "Schedule soft locality", "[scheduler][storage]", - spider::test::StorageTypeList + spider::test::StorageFactoryTypeList ) { - std::tuple< - std::unique_ptr, - std::unique_ptr> - storages = spider::test::create_storage< - std::tuple_element_t<0, TestType>, - std::tuple_element_t<1, TestType>>(); + std::shared_ptr const storage_factory + = spider::test::create_storage_factory(); std::shared_ptr const metadata_store - = std::move(std::get<0>(storages)); - std::shared_ptr const data_store = std::move(std::get<1>(storages)); + = storage_factory->provide_metadata_storage(); + std::shared_ptr const data_store + = storage_factory->provide_data_storage(); - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_store->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + std::shared_ptr const conn + = std::move(std::get>(conn_result)); // Add task boost::uuids::random_generator gen; @@ -170,14 +164,14 @@ TEMPLATE_LIST_TEST_CASE( spider::core::Data data; data.set_hard_locality(false); data.set_locality({"127.0.0.1"}); - REQUIRE(metadata_store->add_driver(conn, spider::core::Driver{client_id}).success()); - REQUIRE(data_store->add_driver_data(conn, client_id, data).success()); + REQUIRE(metadata_store->add_driver(*conn, spider::core::Driver{client_id}).success()); + REQUIRE(data_store->add_driver_data(*conn, client_id, data).success()); task.add_input(spider::core::TaskInput{data.get_id()}); spider::core::TaskGraph graph; graph.add_task(task); graph.add_input_task(task.get_id()); graph.add_output_task(task.get_id()); - REQUIRE(metadata_store->add_job(conn, job_id, client_id, graph).success()); + REQUIRE(metadata_store->add_job(*conn, job_id, client_id, graph).success()); spider::scheduler::FifoPolicy policy{metadata_store, data_store, conn}; // Schedule with wrong address @@ -188,7 +182,7 @@ TEMPLATE_LIST_TEST_CASE( REQUIRE(task_id == task.get_id()); } - REQUIRE(metadata_store->remove_job(conn, job_id).success()); + REQUIRE(metadata_store->remove_job(*conn, job_id).success()); } } // namespace diff --git a/tests/scheduler/test-SchedulerServer.cpp b/tests/scheduler/test-SchedulerServer.cpp index d0905393c..c3df2bf41 100644 --- a/tests/scheduler/test-SchedulerServer.cpp +++ b/tests/scheduler/test-SchedulerServer.cpp @@ -3,7 +3,6 @@ #include #include #include -#include #include #include #include @@ -25,7 +24,8 @@ #include "../../src/spider/scheduler/SchedulerServer.hpp" #include "../../src/spider/storage/DataStorage.hpp" #include "../../src/spider/storage/MetadataStorage.hpp" -#include "../../src/spider/storage/mysql/MySqlConnection.hpp" +#include "../../src/spider/storage/StorageConnection.hpp" +#include "../../src/spider/storage/StorageFactory.hpp" #include "../../src/spider/utils/StopToken.hpp" #include "../storage/StorageTestHelper.hpp" @@ -36,22 +36,20 @@ constexpr int cServerWarmupTime = 5; TEMPLATE_LIST_TEST_CASE( "Scheduler server test", "[scheduler][server][storage]", - spider::test::StorageTypeList + spider::test::StorageFactoryTypeList ) { - std::tuple< - std::unique_ptr, - std::unique_ptr> - storages = spider::test::create_storage< - std::tuple_element_t<0, TestType>, - std::tuple_element_t<1, TestType>>(); + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); std::shared_ptr const metadata_store - = std::move(std::get<0>(storages)); - std::shared_ptr const data_store = std::move(std::get<1>(storages)); + = storage_factory->provide_metadata_storage(); + std::shared_ptr const data_store + = storage_factory->provide_data_storage(); - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_store->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + std::shared_ptr const conn + = std::move(std::get>(conn_result)); std::shared_ptr const policy = std::make_shared(metadata_store, data_store, conn); @@ -84,7 +82,7 @@ TEMPLATE_LIST_TEST_CASE( graph.add_output_task(child_task.get_id()); boost::uuids::random_generator gen; boost::uuids::uuid const job_id = gen(); - REQUIRE(metadata_store->add_job(conn, job_id, gen(), graph).success()); + REQUIRE(metadata_store->add_job(*conn, job_id, gen(), graph).success()); // Schedule request should succeed spider::scheduler::ScheduleTaskRequest const req{gen(), ""}; @@ -99,7 +97,7 @@ TEMPLATE_LIST_TEST_CASE( // Get response should succeed and get child task std::optional const& res_buffer = spider::core::receive_message(socket); - REQUIRE(metadata_store->remove_job(conn, job_id).success()); + REQUIRE(metadata_store->remove_job(*conn, job_id).success()); REQUIRE(res_buffer.has_value()); if (res_buffer.has_value()) { msgpack::object_handle const handle diff --git a/tests/storage/StorageTestHelper.hpp b/tests/storage/StorageTestHelper.hpp index b79795e07..3dd472188 100644 --- a/tests/storage/StorageTestHelper.hpp +++ b/tests/storage/StorageTestHelper.hpp @@ -4,64 +4,28 @@ #include #include +#include #include -#include -#include -#include - -#include "../../src/spider/core/Error.hpp" -#include "../../src/spider/storage/DataStorage.hpp" -#include "../../src/spider/storage/MetadataStorage.hpp" -#include "../../src/spider/storage/mysql/MySqlConnection.hpp" -#include "../../src/spider/storage/mysql/MySqlStorage.hpp" +#include "../../src/spider/storage/mysql/MySqlStorageFactory.hpp" +#include "../../src/spider/storage/StorageFactory.hpp" namespace spider::test { -char const* const cStorageUrl +std::string const cMySqlStorageUrl = "jdbc:mariadb://localhost:3306/spider_test?user=root&password=password"; -using DataStorageTypeList = std::tuple; -using MetadataStorageTypeList = std::tuple; -using StorageTypeList = std::tuple>; +using StorageFactoryTypeList = std::tuple; template -requires std::derived_from -auto create_data_storage() -> std::unique_ptr { - std::unique_ptr storage = std::make_unique(cStorageUrl); - std::variant conn_result - = core::MySqlConnection::create(cStorageUrl); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); - REQUIRE(storage->initialize(conn).success()); - return storage; +requires std::same_as +auto create_storage_factory() -> std::unique_ptr { + return std::make_unique(cMySqlStorageUrl); } template -requires std::derived_from -auto create_metadata_storage() -> std::unique_ptr { - std::unique_ptr storage = std::make_unique(cStorageUrl); - std::variant conn_result - = core::MySqlConnection::create(cStorageUrl); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); - REQUIRE(storage->initialize(conn).success()); - return storage; -} - -template -requires std::derived_from && std::derived_from -auto create_storage( -) -> std::tuple, std::unique_ptr> { - std::variant conn_result - = core::MySqlConnection::create(cStorageUrl); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); - - std::unique_ptr metadata_storage = std::make_unique(cStorageUrl); - REQUIRE(metadata_storage->initialize(conn).success()); - std::unique_ptr data_storage = std::make_unique(cStorageUrl); - REQUIRE(data_storage->initialize(conn).success()); - return std::make_tuple(std::move(metadata_storage), std::move(data_storage)); +requires std::same_as +auto get_storage_url() -> std::string { + return cMySqlStorageUrl; } } // namespace spider::test diff --git a/tests/storage/test-DataStorage.cpp b/tests/storage/test-DataStorage.cpp index 1e4a196c3..21da3bb62 100644 --- a/tests/storage/test-DataStorage.cpp +++ b/tests/storage/test-DataStorage.cpp @@ -1,5 +1,6 @@ // NOLINTBEGIN(cert-err58-cpp,cppcoreguidelines-avoid-do-while,readability-function-cognitive-complexity,cppcoreguidelines-avoid-non-const-global-variables,cppcoreguidelines-avoid-c-arrays,modernize-avoid-c-arrays) -#include +#include +#include #include #include @@ -13,92 +14,111 @@ #include "../../src/spider/core/KeyValueData.hpp" #include "../../src/spider/core/Task.hpp" #include "../../src/spider/core/TaskGraph.hpp" -#include "../../src/spider/storage/mysql/MySqlConnection.hpp" +#include "../../src/spider/storage/DataStorage.hpp" +#include "../../src/spider/storage/MetadataStorage.hpp" +#include "../../src/spider/storage/StorageConnection.hpp" +#include "../../src/spider/storage/StorageFactory.hpp" #include "../utils/CoreDataUtils.hpp" #include "StorageTestHelper.hpp" namespace { -TEMPLATE_LIST_TEST_CASE("Add, get and remove data", "[storage]", spider::test::StorageTypeList) { - auto [metadata_storage, data_storage] = spider::test:: - create_storage, std::tuple_element_t<1, TestType>>(); - - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_storage->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); +TEMPLATE_LIST_TEST_CASE( + "Add, get and remove data", + "[storage]", + spider::test::StorageFactoryTypeList +) { + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); + std::unique_ptr metadata_storage + = storage_factory->provide_metadata_storage(); + std::unique_ptr data_storage + = storage_factory->provide_data_storage(); + + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + auto conn = std::move(std::get>(conn_result)); // Add driver and data spider::core::Data const data{"value"}; boost::uuids::random_generator gen; boost::uuids::uuid const driver_id = gen(); - REQUIRE(metadata_storage->add_driver(conn, spider::core::Driver{driver_id}).success()); - REQUIRE(data_storage->add_driver_data(conn, driver_id, data).success()); + REQUIRE(metadata_storage->add_driver(*conn, spider::core::Driver{driver_id}).success()); + REQUIRE(data_storage->add_driver_data(*conn, driver_id, data).success()); // Add data with same id again should fail spider::core::Data const data_same_id{data.get_id(), "value2"}; REQUIRE(spider::core::StorageErrType::DuplicateKeyErr - == data_storage->add_driver_data(conn, driver_id, data_same_id).type); + == data_storage->add_driver_data(*conn, driver_id, data_same_id).type); // Get data should match spider::core::Data result{"temp"}; - REQUIRE(data_storage->get_data(conn, data.get_id(), &result).success()); + REQUIRE(data_storage->get_data(*conn, data.get_id(), &result).success()); REQUIRE(spider::test::data_equal(data, result)); // Remove data should succeed - REQUIRE(data_storage->remove_data(conn, data.get_id()).success()); + REQUIRE(data_storage->remove_data(*conn, data.get_id()).success()); // Get data should fail REQUIRE(spider::core::StorageErrType::KeyNotFoundErr - == data_storage->get_data(conn, data.get_id(), &result).type); + == data_storage->get_data(*conn, data.get_id(), &result).type); } TEMPLATE_LIST_TEST_CASE( "Add and get driver key value data", "[storage]", - spider::test::StorageTypeList + spider::test::StorageFactoryTypeList ) { - auto [metadata_storage, data_storage] = spider::test:: - create_storage, std::tuple_element_t<1, TestType>>(); - - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_storage->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); + std::unique_ptr metadata_storage + = storage_factory->provide_metadata_storage(); + std::unique_ptr data_storage + = storage_factory->provide_data_storage(); + + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + auto conn = std::move(std::get>(conn_result)); // Add driver boost::uuids::random_generator gen; boost::uuids::uuid const driver_id = gen(); - REQUIRE(metadata_storage->add_driver(conn, spider::core::Driver{driver_id}).success()); + REQUIRE(metadata_storage->add_driver(*conn, spider::core::Driver{driver_id}).success()); // Add data spider::core::KeyValueData const data{"key", "value", driver_id}; - REQUIRE(data_storage->add_client_kv_data(conn, data).success()); + REQUIRE(data_storage->add_client_kv_data(*conn, data).success()); // Add data with same key and id again should fail spider::core::KeyValueData const data_same_key{"key", "value2", driver_id}; REQUIRE(spider::core::StorageErrType::DuplicateKeyErr - == data_storage->add_client_kv_data(conn, data_same_key).type); + == data_storage->add_client_kv_data(*conn, data_same_key).type); // Get data should match std::string value; - auto err = data_storage->get_client_kv_data(conn, driver_id, "key", &value); - REQUIRE(data_storage->get_client_kv_data(conn, driver_id, "key", &value).success()); + auto err = data_storage->get_client_kv_data(*conn, driver_id, "key", &value); + REQUIRE(data_storage->get_client_kv_data(*conn, driver_id, "key", &value).success()); REQUIRE(data.get_value() == value); } TEMPLATE_LIST_TEST_CASE( "Add and get task key value data", "[storage]", - spider::test::StorageTypeList + spider::test::StorageFactoryTypeList ) { - auto [metadata_storage, data_storage] = spider::test:: - create_storage, std::tuple_element_t<1, TestType>>(); - - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_storage->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); + std::unique_ptr metadata_storage + = storage_factory->provide_metadata_storage(); + std::unique_ptr data_storage + = storage_factory->provide_data_storage(); + + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + auto conn = std::move(std::get>(conn_result)); // Add task boost::uuids::random_generator gen; @@ -108,42 +128,46 @@ TEMPLATE_LIST_TEST_CASE( graph.add_input_task(task.get_id()); graph.add_output_task(task.get_id()); boost::uuids::uuid const job_id = gen(); - REQUIRE(metadata_storage->add_job(conn, job_id, gen(), graph).success()); + REQUIRE(metadata_storage->add_job(*conn, job_id, gen(), graph).success()); // Add data spider::core::KeyValueData const data{"key", "value", task.get_id()}; - REQUIRE(data_storage->add_task_kv_data(conn, data).success()); + REQUIRE(data_storage->add_task_kv_data(*conn, data).success()); // Add data with same key and id again should fail spider::core::KeyValueData const data_same_key{"key", "value2", task.get_id()}; REQUIRE(spider::core::StorageErrType::DuplicateKeyErr - == data_storage->add_task_kv_data(conn, data_same_key).type); + == data_storage->add_task_kv_data(*conn, data_same_key).type); // Get data should match std::string value; - REQUIRE(data_storage->get_task_kv_data(conn, task.get_id(), "key", &value).success()); + REQUIRE(data_storage->get_task_kv_data(*conn, task.get_id(), "key", &value).success()); REQUIRE(data.get_value() == value); // Clean up - REQUIRE(metadata_storage->remove_job(conn, job_id).success()); + REQUIRE(metadata_storage->remove_job(*conn, job_id).success()); } TEMPLATE_LIST_TEST_CASE( "Add and remove task reference for task", "[storage]", - spider::test::StorageTypeList + spider::test::StorageFactoryTypeList ) { - auto [metadata_storage, data_storage] = spider::test:: - create_storage, std::tuple_element_t<1, TestType>>(); - - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_storage->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); + std::unique_ptr metadata_storage + = storage_factory->provide_metadata_storage(); + std::unique_ptr data_storage + = storage_factory->provide_data_storage(); + + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + auto conn = std::move(std::get>(conn_result)); boost::uuids::random_generator gen; // Add task reference without data and task should fail. - REQUIRE(!data_storage->add_task_reference(conn, gen(), gen()).success()); + REQUIRE(!data_storage->add_task_reference(*conn, gen(), gen()).success()); // Add task spider::core::Task const task{"func"}; @@ -155,69 +179,73 @@ TEMPLATE_LIST_TEST_CASE( graph.add_input_task(task.get_id()); graph.add_output_task(task_2.get_id()); boost::uuids::uuid const job_id = gen(); - REQUIRE(metadata_storage->add_job(conn, job_id, gen(), graph).success()); + REQUIRE(metadata_storage->add_job(*conn, job_id, gen(), graph).success()); // Add task reference without data should fail. - REQUIRE(!data_storage->add_task_reference(conn, gen(), task.get_id()).success()); + REQUIRE(!data_storage->add_task_reference(*conn, gen(), task.get_id()).success()); // Add data spider::core::Data const data{"value"}; - REQUIRE(data_storage->add_task_data(conn, task.get_id(), data).success()); + REQUIRE(data_storage->add_task_data(*conn, task.get_id(), data).success()); // Add task reference - REQUIRE(data_storage->add_task_reference(conn, data.get_id(), task_2.get_id()).success()); + REQUIRE(data_storage->add_task_reference(*conn, data.get_id(), task_2.get_id()).success()); // Remove task reference - REQUIRE(data_storage->remove_task_reference(conn, data.get_id(), task_2.get_id()).success()); + REQUIRE(data_storage->remove_task_reference(*conn, data.get_id(), task_2.get_id()).success()); // Remove job - REQUIRE(metadata_storage->remove_job(conn, job_id).success()); + REQUIRE(metadata_storage->remove_job(*conn, job_id).success()); // Clean up - REQUIRE(data_storage->remove_dangling_data(conn).success()); + REQUIRE(data_storage->remove_dangling_data(*conn).success()); // Get data should fail spider::core::Data res{"temp"}; REQUIRE(spider::core::StorageErrType::KeyNotFoundErr - == data_storage->get_data(conn, data.get_id(), &res).type); + == data_storage->get_data(*conn, data.get_id(), &res).type); } TEMPLATE_LIST_TEST_CASE( "Add and remove data reference for driver", "[storage]", - spider::test::StorageTypeList + spider::test::StorageFactoryTypeList ) { - auto [metadata_storage, data_storage] = spider::test:: - create_storage, std::tuple_element_t<1, TestType>>(); - - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_storage->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); + std::unique_ptr metadata_storage + = storage_factory->provide_metadata_storage(); + std::unique_ptr data_storage + = storage_factory->provide_data_storage(); + + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + auto conn = std::move(std::get>(conn_result)); boost::uuids::random_generator gen; // Add driver reference without data and driver should fail - REQUIRE(!data_storage->add_driver_reference(conn, gen(), gen()).success()); + REQUIRE(!data_storage->add_driver_reference(*conn, gen(), gen()).success()); // Add driver boost::uuids::uuid const driver_id = gen(); boost::uuids::uuid const driver_id_2 = gen(); - REQUIRE(metadata_storage->add_driver(conn, spider::core::Driver{driver_id}).success()); - REQUIRE(metadata_storage->add_driver(conn, spider::core::Driver{driver_id_2}).success()); + REQUIRE(metadata_storage->add_driver(*conn, spider::core::Driver{driver_id}).success()); + REQUIRE(metadata_storage->add_driver(*conn, spider::core::Driver{driver_id_2}).success()); // Add driver reference without data should fail - REQUIRE(!data_storage->add_driver_reference(conn, gen(), driver_id).success()); + REQUIRE(!data_storage->add_driver_reference(*conn, gen(), driver_id).success()); // Add data spider::core::Data const data{"value"}; - REQUIRE(data_storage->add_driver_data(conn, driver_id, data).success()); + REQUIRE(data_storage->add_driver_data(*conn, driver_id, data).success()); // Add driver reference - REQUIRE(data_storage->add_driver_reference(conn, data.get_id(), driver_id_2).success()); + REQUIRE(data_storage->add_driver_reference(*conn, data.get_id(), driver_id_2).success()); // Remove driver reference - REQUIRE(data_storage->remove_driver_reference(conn, data.get_id(), driver_id_2).success()); + REQUIRE(data_storage->remove_driver_reference(*conn, data.get_id(), driver_id_2).success()); } } // namespace diff --git a/tests/storage/test-MetadataStorage.cpp b/tests/storage/test-MetadataStorage.cpp index 9cc4612fc..22b8b2003 100644 --- a/tests/storage/test-MetadataStorage.cpp +++ b/tests/storage/test-MetadataStorage.cpp @@ -4,6 +4,7 @@ #include #include #include +#include #include #include @@ -17,33 +18,36 @@ #include "../../src/spider/core/JobMetadata.hpp" #include "../../src/spider/core/Task.hpp" #include "../../src/spider/core/TaskGraph.hpp" +#include "../../src/spider/storage/JobSubmissionBatch.hpp" #include "../../src/spider/storage/MetadataStorage.hpp" -#include "../../src/spider/storage/mysql/MySqlConnection.hpp" -#include "../../src/spider/storage/mysql/MySqlJobSubmissionBatch.hpp" +#include "../../src/spider/storage/StorageConnection.hpp" +#include "../../src/spider/storage/StorageFactory.hpp" #include "../utils/CoreTaskUtils.hpp" #include "StorageTestHelper.hpp" namespace { -TEMPLATE_LIST_TEST_CASE("Driver heartbeat", "[storage]", spider::test::MetadataStorageTypeList) { +TEMPLATE_LIST_TEST_CASE("Driver heartbeat", "[storage]", spider::test::StorageFactoryTypeList) { + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); std::unique_ptr storage - = spider::test::create_metadata_storage(); + = storage_factory->provide_metadata_storage(); - std::variant conn_result - = spider::core::MySqlConnection::create(storage->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + auto conn = std::move(std::get>(conn_result)); constexpr double cDuration = 100; // Add driver should succeed boost::uuids::random_generator gen; boost::uuids::uuid const driver_id = gen(); - REQUIRE(storage->add_driver(conn, spider::core::Driver{driver_id}).success()); + REQUIRE(storage->add_driver(*conn, spider::core::Driver{driver_id}).success()); std::vector ids{}; // Driver should not time out - REQUIRE(storage->heartbeat_timeout(conn, cDuration, &ids).success()); + REQUIRE(storage->heartbeat_timeout(*conn, cDuration, &ids).success()); // Because other tests may run in parallel, just check `ids` don't have `driver_id` REQUIRE(std::ranges::none_of(ids, [&driver_id](boost::uuids::uuid id) { return id == driver_id; @@ -52,7 +56,7 @@ TEMPLATE_LIST_TEST_CASE("Driver heartbeat", "[storage]", spider::test::MetadataS std::this_thread::sleep_for(std::chrono::seconds(1)); // Driver should time out - REQUIRE(storage->heartbeat_timeout(conn, cDuration, &ids).success()); + REQUIRE(storage->heartbeat_timeout(*conn, cDuration, &ids).success()); REQUIRE(!ids.empty()); REQUIRE(std::ranges::any_of(ids, [&driver_id](boost::uuids::uuid id) { return id == driver_id; @@ -60,9 +64,9 @@ TEMPLATE_LIST_TEST_CASE("Driver heartbeat", "[storage]", spider::test::MetadataS ids.clear(); // Update heartbeat - REQUIRE(storage->update_heartbeat(conn, driver_id).success()); + REQUIRE(storage->update_heartbeat(*conn, driver_id).success()); // Driver should not time out - REQUIRE(storage->heartbeat_timeout(conn, cDuration, &ids).success()); + REQUIRE(storage->heartbeat_timeout(*conn, cDuration, &ids).success()); REQUIRE(std::ranges::none_of(ids, [&driver_id](boost::uuids::uuid id) { return id == driver_id; })); @@ -71,63 +75,69 @@ TEMPLATE_LIST_TEST_CASE("Driver heartbeat", "[storage]", spider::test::MetadataS TEMPLATE_LIST_TEST_CASE( "Scheduler state and addr", "[storage]", - spider::test::MetadataStorageTypeList + spider::test::StorageFactoryTypeList ) { + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); std::unique_ptr storage - = spider::test::create_metadata_storage(); + = storage_factory->provide_metadata_storage(); - std::variant conn_result - = spider::core::MySqlConnection::create(storage->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + auto conn = std::move(std::get>(conn_result)); boost::uuids::random_generator gen; boost::uuids::uuid const scheduler_id = gen(); constexpr int cPort = 3306; // Add scheduler should succeed - REQUIRE(storage->add_scheduler(conn, spider::core::Scheduler{scheduler_id, "127.0.0.1", cPort}) + REQUIRE(storage->add_scheduler(*conn, spider::core::Scheduler{scheduler_id, "127.0.0.1", cPort}) .success()); // Get scheduler addr should succeed std::string addr_res; int port_res = 0; - REQUIRE(storage->get_scheduler_addr(conn, scheduler_id, &addr_res, &port_res).success()); + REQUIRE(storage->get_scheduler_addr(*conn, scheduler_id, &addr_res, &port_res).success()); REQUIRE(addr_res == "127.0.0.1"); REQUIRE(port_res == cPort); // Get non-exist scheduler should fail REQUIRE(spider::core::StorageErrType::KeyNotFoundErr - == storage->get_scheduler_addr(conn, gen(), &addr_res, &port_res).type); + == storage->get_scheduler_addr(*conn, gen(), &addr_res, &port_res).type); // Get default state std::string state_res; - REQUIRE(storage->get_scheduler_state(conn, scheduler_id, &state_res).success()); + REQUIRE(storage->get_scheduler_state(*conn, scheduler_id, &state_res).success()); REQUIRE(state_res == "normal"); state_res.clear(); // Update scheduler state should succeed std::string state = "recovery"; - REQUIRE(storage->set_scheduler_state(conn, scheduler_id, state).success()); + REQUIRE(storage->set_scheduler_state(*conn, scheduler_id, state).success()); // Get new state - REQUIRE(storage->get_scheduler_state(conn, scheduler_id, &state_res).success()); + REQUIRE(storage->get_scheduler_state(*conn, scheduler_id, &state_res).success()); REQUIRE(state_res == state); } TEMPLATE_LIST_TEST_CASE( - "Job add, get and remove", + "Job batch add, get and remove", "[storage]", - spider::test::MetadataStorageTypeList + spider::test::StorageFactoryTypeList ) { + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); std::unique_ptr storage - = spider::test::create_metadata_storage(); + = storage_factory->provide_metadata_storage(); + + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + auto conn = std::move(std::get>(conn_result)); - std::variant conn_result - = spider::core::MySqlConnection::create(storage->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); - spider::core::MySqlJobSubmissionBatch batch{conn}; + std::unique_ptr batch + = storage_factory->provide_job_submission_batch(*conn); boost::uuids::random_generator gen; boost::uuids::uuid const job_id = gen(); @@ -179,17 +189,18 @@ TEMPLATE_LIST_TEST_CASE( REQUIRE(heads[0] == simple_task.get_id()); // Submit job should success - REQUIRE(storage->add_job_batch(conn, batch, job_id, client_id, graph).success()); - REQUIRE(storage->add_job_batch(conn, batch, simple_job_id, client_id, simple_graph).success()); - batch.submit_batch(conn); + REQUIRE(storage->add_job_batch(*conn, *batch, job_id, client_id, graph).success()); + REQUIRE(storage->add_job_batch(*conn, *batch, simple_job_id, client_id, simple_graph).success() + ); + batch->submit_batch(*conn); // Get job id for non-existent client id should return empty vector std::vector job_ids; - REQUIRE(storage->get_jobs_by_client_id(conn, gen(), &job_ids).success()); + REQUIRE(storage->get_jobs_by_client_id(*conn, gen(), &job_ids).success()); REQUIRE(job_ids.empty()); // Get job id for client id should get correct value - REQUIRE(storage->get_jobs_by_client_id(conn, client_id, &job_ids).success()); + REQUIRE(storage->get_jobs_by_client_id(*conn, client_id, &job_ids).success()); REQUIRE(2 == job_ids.size()); REQUIRE( ((job_ids[0] == job_id && job_ids[1] == simple_job_id) @@ -198,7 +209,7 @@ TEMPLATE_LIST_TEST_CASE( // Get job metadata should get correct value spider::core::JobMetadata job_metadata{}; - REQUIRE(storage->get_job_metadata(conn, job_id, &job_metadata).success()); + REQUIRE(storage->get_job_metadata(*conn, job_id, &job_metadata).success()); REQUIRE(job_id == job_metadata.get_id()); REQUIRE(client_id == job_metadata.get_client_id()); std::chrono::seconds const time_delta{1}; @@ -207,26 +218,26 @@ TEMPLATE_LIST_TEST_CASE( // Get task graph should succeed spider::core::TaskGraph graph_res{}; - REQUIRE(storage->get_task_graph(conn, job_id, &graph_res).success()); + REQUIRE(storage->get_task_graph(*conn, job_id, &graph_res).success()); REQUIRE(spider::test::task_graph_equal(graph, graph_res)); spider::core::TaskGraph simple_graph_res{}; - REQUIRE(storage->get_task_graph(conn, simple_job_id, &simple_graph_res).success()); + REQUIRE(storage->get_task_graph(*conn, simple_job_id, &simple_graph_res).success()); REQUIRE(spider::test::task_graph_equal(simple_graph, simple_graph_res)); // Get task should succeed spider::core::Task task_res{""}; - REQUIRE(storage->get_task(conn, child_task.get_id(), &task_res).success()); + REQUIRE(storage->get_task(*conn, child_task.get_id(), &task_res).success()); REQUIRE(spider::test::task_equal(child_task, task_res)); // Get child tasks should succeed std::vector tasks; - REQUIRE(storage->get_child_tasks(conn, parent_1.get_id(), &tasks).success()); + REQUIRE(storage->get_child_tasks(*conn, parent_1.get_id(), &tasks).success()); REQUIRE(1 == tasks.size()); REQUIRE(spider::test::task_equal(child_task, tasks[0])); tasks.clear(); // Get parent tasks should succeed - REQUIRE(storage->get_parent_tasks(conn, child_task.get_id(), &tasks).success()); + REQUIRE(storage->get_parent_tasks(*conn, child_task.get_id(), &tasks).success()); REQUIRE(2 == tasks.size()); REQUIRE( ((spider::test::task_equal(tasks[0], parent_1) @@ -236,27 +247,29 @@ TEMPLATE_LIST_TEST_CASE( ); // Remove job should succeed - REQUIRE(storage->remove_job(conn, simple_job_id).success()); + REQUIRE(storage->remove_job(*conn, simple_job_id).success()); REQUIRE(spider::core::StorageErrType::KeyNotFoundErr - == storage->get_task_graph(conn, simple_job_id, &simple_graph_res).type); + == storage->get_task_graph(*conn, simple_job_id, &simple_graph_res).type); graph_res = spider::core::TaskGraph{}; - REQUIRE(storage->get_task_graph(conn, job_id, &graph_res).success()); + REQUIRE(storage->get_task_graph(*conn, job_id, &graph_res).success()); REQUIRE(spider::test::task_graph_equal(graph, graph_res)); - REQUIRE(storage->remove_job(conn, job_id).success()); + REQUIRE(storage->remove_job(*conn, job_id).success()); } TEMPLATE_LIST_TEST_CASE( - "Job batch add, get and remove", + "Job add, get and remove", "[storage]", - spider::test::MetadataStorageTypeList + spider::test::StorageFactoryTypeList ) { + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); std::unique_ptr storage - = spider::test::create_metadata_storage(); + = storage_factory->provide_metadata_storage(); - std::variant conn_result - = spider::core::MySqlConnection::create(storage->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + auto conn = std::move(std::get>(conn_result)); boost::uuids::random_generator gen; boost::uuids::uuid const job_id = gen(); @@ -308,16 +321,16 @@ TEMPLATE_LIST_TEST_CASE( REQUIRE(heads[0] == simple_task.get_id()); // Submit job should success - REQUIRE(storage->add_job(conn, job_id, client_id, graph).success()); - REQUIRE(storage->add_job(conn, simple_job_id, client_id, simple_graph).success()); + REQUIRE(storage->add_job(*conn, job_id, client_id, graph).success()); + REQUIRE(storage->add_job(*conn, simple_job_id, client_id, simple_graph).success()); // Get job id for non-existent client id should return empty vector std::vector job_ids; - REQUIRE(storage->get_jobs_by_client_id(conn, gen(), &job_ids).success()); + REQUIRE(storage->get_jobs_by_client_id(*conn, gen(), &job_ids).success()); REQUIRE(job_ids.empty()); // Get job id for client id should get correct value - REQUIRE(storage->get_jobs_by_client_id(conn, client_id, &job_ids).success()); + REQUIRE(storage->get_jobs_by_client_id(*conn, client_id, &job_ids).success()); REQUIRE(2 == job_ids.size()); REQUIRE( ((job_ids[0] == job_id && job_ids[1] == simple_job_id) @@ -326,7 +339,7 @@ TEMPLATE_LIST_TEST_CASE( // Get job metadata should get correct value spider::core::JobMetadata job_metadata{}; - REQUIRE(storage->get_job_metadata(conn, job_id, &job_metadata).success()); + REQUIRE(storage->get_job_metadata(*conn, job_id, &job_metadata).success()); REQUIRE(job_id == job_metadata.get_id()); REQUIRE(client_id == job_metadata.get_client_id()); std::chrono::seconds const time_delta{1}; @@ -335,26 +348,26 @@ TEMPLATE_LIST_TEST_CASE( // Get task graph should succeed spider::core::TaskGraph graph_res{}; - REQUIRE(storage->get_task_graph(conn, job_id, &graph_res).success()); + REQUIRE(storage->get_task_graph(*conn, job_id, &graph_res).success()); REQUIRE(spider::test::task_graph_equal(graph, graph_res)); spider::core::TaskGraph simple_graph_res{}; - REQUIRE(storage->get_task_graph(conn, simple_job_id, &simple_graph_res).success()); + REQUIRE(storage->get_task_graph(*conn, simple_job_id, &simple_graph_res).success()); REQUIRE(spider::test::task_graph_equal(simple_graph, simple_graph_res)); // Get task should succeed spider::core::Task task_res{""}; - REQUIRE(storage->get_task(conn, child_task.get_id(), &task_res).success()); + REQUIRE(storage->get_task(*conn, child_task.get_id(), &task_res).success()); REQUIRE(spider::test::task_equal(child_task, task_res)); // Get child tasks should succeed std::vector tasks; - REQUIRE(storage->get_child_tasks(conn, parent_1.get_id(), &tasks).success()); + REQUIRE(storage->get_child_tasks(*conn, parent_1.get_id(), &tasks).success()); REQUIRE(1 == tasks.size()); REQUIRE(spider::test::task_equal(child_task, tasks[0])); tasks.clear(); // Get parent tasks should succeed - REQUIRE(storage->get_parent_tasks(conn, child_task.get_id(), &tasks).success()); + REQUIRE(storage->get_parent_tasks(*conn, child_task.get_id(), &tasks).success()); REQUIRE(2 == tasks.size()); REQUIRE( ((spider::test::task_equal(tasks[0], parent_1) @@ -364,23 +377,25 @@ TEMPLATE_LIST_TEST_CASE( ); // Remove job should succeed - REQUIRE(storage->remove_job(conn, simple_job_id).success()); + REQUIRE(storage->remove_job(*conn, simple_job_id).success()); REQUIRE(spider::core::StorageErrType::KeyNotFoundErr - == storage->get_task_graph(conn, simple_job_id, &simple_graph_res).type); + == storage->get_task_graph(*conn, simple_job_id, &simple_graph_res).type); graph_res = spider::core::TaskGraph{}; - REQUIRE(storage->get_task_graph(conn, job_id, &graph_res).success()); + REQUIRE(storage->get_task_graph(*conn, job_id, &graph_res).success()); REQUIRE(spider::test::task_graph_equal(graph, graph_res)); - REQUIRE(storage->remove_job(conn, job_id).success()); + REQUIRE(storage->remove_job(*conn, job_id).success()); } -TEMPLATE_LIST_TEST_CASE("Task finish", "[storage]", spider::test::MetadataStorageTypeList) { +TEMPLATE_LIST_TEST_CASE("Task finish", "[storage]", spider::test::StorageFactoryTypeList) { + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); std::unique_ptr storage - = spider::test::create_metadata_storage(); + = storage_factory->provide_metadata_storage(); - std::variant conn_result - = spider::core::MySqlConnection::create(storage->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + auto conn = std::move(std::get>(conn_result)); boost::uuids::random_generator gen; boost::uuids::uuid const job_id = gen(); @@ -409,49 +424,51 @@ TEMPLATE_LIST_TEST_CASE("Task finish", "[storage]", spider::test::MetadataStorag graph.add_input_task(parent_2.get_id()); graph.add_output_task(child_task.get_id()); // Submit job should success - REQUIRE(storage->add_job(conn, job_id, gen(), graph).success()); + REQUIRE(storage->add_job(*conn, job_id, gen(), graph).success()); // Task finish for parent 1 should succeed spider::core::TaskInstance const parent_1_instance{gen(), parent_1.get_id()}; - REQUIRE(storage->set_task_state(conn, parent_1.get_id(), spider::core::TaskState::Running) + REQUIRE(storage->set_task_state(*conn, parent_1.get_id(), spider::core::TaskState::Running) .success()); REQUIRE(storage->task_finish( - conn, + *conn, parent_1_instance, {spider::core::TaskOutput{"1.1", "float"}} ).success()); // Parent 1 finish should not update state of any other tasks spider::core::Task res_task{""}; - REQUIRE(storage->get_task(conn, parent_2.get_id(), &res_task).success()); + REQUIRE(storage->get_task(*conn, parent_2.get_id(), &res_task).success()); REQUIRE(spider::test::task_equal(parent_2, res_task)); REQUIRE(res_task.get_state() == spider::core::TaskState::Ready); - REQUIRE(storage->get_task(conn, child_task.get_id(), &res_task).success()); + REQUIRE(storage->get_task(*conn, child_task.get_id(), &res_task).success()); REQUIRE(res_task.get_state() == spider::core::TaskState::Pending); // Task finish for parent 2 should success spider::core::TaskInstance const parent_2_instance{gen(), parent_2.get_id()}; - REQUIRE(storage->set_task_state(conn, parent_2.get_id(), spider::core::TaskState::Running) + REQUIRE(storage->set_task_state(*conn, parent_2.get_id(), spider::core::TaskState::Running) .success()); - REQUIRE(storage->task_finish(conn, parent_2_instance, {spider::core::TaskOutput{"2", "int"}}) + REQUIRE(storage->task_finish(*conn, parent_2_instance, {spider::core::TaskOutput{"2", "int"}}) .success()); // Parent 2 finish should update state of child - REQUIRE(storage->get_task(conn, child_task.get_id(), &res_task).success()); + REQUIRE(storage->get_task(*conn, child_task.get_id(), &res_task).success()); REQUIRE(res_task.get_input(0).get_value() == "1.1"); REQUIRE(res_task.get_input(1).get_value() == "2"); REQUIRE(res_task.get_state() == spider::core::TaskState::Ready); // Clean up - REQUIRE(storage->remove_job(conn, job_id).success()); + REQUIRE(storage->remove_job(*conn, job_id).success()); } -TEMPLATE_LIST_TEST_CASE("Job reset", "[storage]", spider::test::MetadataStorageTypeList) { +TEMPLATE_LIST_TEST_CASE("Job reset", "[storage]", spider::test::StorageFactoryTypeList) { + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); std::unique_ptr storage - = spider::test::create_metadata_storage(); + = storage_factory->provide_metadata_storage(); - std::variant conn_result - = spider::core::MySqlConnection::create(storage->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + auto conn = std::move(std::get>(conn_result)); boost::uuids::random_generator gen; boost::uuids::uuid const job_id = gen(); @@ -483,51 +500,51 @@ TEMPLATE_LIST_TEST_CASE("Job reset", "[storage]", spider::test::MetadataStorageT graph.add_input_task(parent_2.get_id()); graph.add_output_task(child_task.get_id()); // Submit job should success - REQUIRE(storage->add_job(conn, job_id, gen(), graph).success()); + REQUIRE(storage->add_job(*conn, job_id, gen(), graph).success()); // Task finish for parent 1 should succeed spider::core::TaskInstance const parent_1_instance{gen(), parent_1.get_id()}; - REQUIRE(storage->set_task_state(conn, parent_1.get_id(), spider::core::TaskState::Running) + REQUIRE(storage->set_task_state(*conn, parent_1.get_id(), spider::core::TaskState::Running) .success()); REQUIRE(storage->task_finish( - conn, + *conn, parent_1_instance, {spider::core::TaskOutput{"1.1", "float"}} ).success()); // Task finish for parent 2 should success spider::core::TaskInstance const parent_2_instance{gen(), parent_2.get_id()}; - REQUIRE(storage->set_task_state(conn, parent_2.get_id(), spider::core::TaskState::Running) + REQUIRE(storage->set_task_state(*conn, parent_2.get_id(), spider::core::TaskState::Running) .success()); - REQUIRE(storage->task_finish(conn, parent_2_instance, {spider::core::TaskOutput{"2", "int"}}) + REQUIRE(storage->task_finish(*conn, parent_2_instance, {spider::core::TaskOutput{"2", "int"}}) .success()); // Task finish for child should success spider::core::TaskInstance const child_instance{gen(), child_task.get_id()}; - REQUIRE(storage->set_task_state(conn, child_task.get_id(), spider::core::TaskState::Running) + REQUIRE(storage->set_task_state(*conn, child_task.get_id(), spider::core::TaskState::Running) .success()); - REQUIRE(storage->task_finish(conn, child_instance, {spider::core::TaskOutput{"3.3", "float"}}) + REQUIRE(storage->task_finish(*conn, child_instance, {spider::core::TaskOutput{"3.3", "float"}}) .success()); // Job reset - REQUIRE(storage->reset_job(conn, job_id).success()); + REQUIRE(storage->reset_job(*conn, job_id).success()); // Parent tasks states should be ready and child task state should be waiting // Parent tasks inputs should be available and child task inputs should be empty // All tasks output should be empty spider::core::Task res_task{""}; - REQUIRE(storage->get_task(conn, parent_1.get_id(), &res_task).success()); + REQUIRE(storage->get_task(*conn, parent_1.get_id(), &res_task).success()); REQUIRE(res_task.get_state() == spider::core::TaskState::Ready); REQUIRE(res_task.get_num_inputs() == 2); REQUIRE(res_task.get_input(0).get_value() == "1"); REQUIRE(res_task.get_input(1).get_value() == "2"); REQUIRE(res_task.get_num_outputs() == 1); REQUIRE(!res_task.get_output(0).get_value().has_value()); - REQUIRE(storage->get_task(conn, parent_2.get_id(), &res_task).success()); + REQUIRE(storage->get_task(*conn, parent_2.get_id(), &res_task).success()); REQUIRE(res_task.get_state() == spider::core::TaskState::Ready); REQUIRE(res_task.get_num_inputs() == 2); REQUIRE(res_task.get_input(0).get_value() == "3"); REQUIRE(res_task.get_input(1).get_value() == "4"); REQUIRE(res_task.get_num_outputs() == 1); REQUIRE(!res_task.get_output(0).get_value().has_value()); - REQUIRE(storage->get_task(conn, child_task.get_id(), &res_task).success()); + REQUIRE(storage->get_task(*conn, child_task.get_id(), &res_task).success()); REQUIRE(res_task.get_state() == spider::core::TaskState::Pending); REQUIRE(res_task.get_num_inputs() == 2); REQUIRE(!res_task.get_input(0).get_value().has_value()); @@ -536,7 +553,7 @@ TEMPLATE_LIST_TEST_CASE("Job reset", "[storage]", spider::test::MetadataStorageT REQUIRE(!res_task.get_output(0).get_value().has_value()); // Clean up - REQUIRE(storage->remove_job(conn, job_id).success()); + REQUIRE(storage->remove_job(*conn, job_id).success()); } } // namespace diff --git a/tests/worker/test-FunctionManager.cpp b/tests/worker/test-FunctionManager.cpp index 5f889bbde..de59c9730 100644 --- a/tests/worker/test-FunctionManager.cpp +++ b/tests/worker/test-FunctionManager.cpp @@ -17,7 +17,10 @@ #include "../../src/spider/core/Error.hpp" #include "../../src/spider/core/TaskContextImpl.hpp" #include "../../src/spider/io/MsgPack.hpp" // IWYU pragma: keep -#include "../../src/spider/storage/mysql/MySqlConnection.hpp" +#include "../../src/spider/storage/DataStorage.hpp" +#include "../../src/spider/storage/MetadataStorage.hpp" +#include "../../src/spider/storage/StorageConnection.hpp" +#include "../../src/spider/storage/StorageFactory.hpp" #include "../../src/spider/worker/FunctionManager.hpp" #include "../../src/spider/worker/FunctionNameManager.hpp" #include "../storage/StorageTestHelper.hpp" @@ -63,15 +66,21 @@ TEST_CASE("Register and get function name", "[core]") { TEMPLATE_LIST_TEST_CASE( "Register and run function with POD inputs", "[core][storage]", - spider::test::StorageTypeList + spider::test::StorageFactoryTypeList ) { - auto [metadata_storage, data_storage] = spider::test:: - create_storage, std::tuple_element_t<1, TestType>>(); + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); + std::unique_ptr metadata_storage + = storage_factory->provide_metadata_storage(); + std::unique_ptr data_storage + = storage_factory->provide_data_storage(); + boost::uuids::random_generator gen; spider::TaskContext context = spider::core::TaskContextImpl::create_task_context( gen(), std::move(data_storage), - std::move(metadata_storage) + std::move(metadata_storage), + std::move(storage_factory) ); spider::core::FunctionManager const& manager = spider::core::FunctionManager::get_instance(); @@ -116,15 +125,21 @@ TEMPLATE_LIST_TEST_CASE( TEMPLATE_LIST_TEST_CASE( "Register and run function with tuple return", "[core][storage]", - spider::test::StorageTypeList + spider::test::StorageFactoryTypeList ) { - auto [metadata_storage, data_storage] = spider::test:: - create_storage, std::tuple_element_t<1, TestType>>(); + std::unique_ptr storage_factory + = spider::test::create_storage_factory(); + std::unique_ptr metadata_storage + = storage_factory->provide_metadata_storage(); + std::unique_ptr data_storage + = storage_factory->provide_data_storage(); + boost::uuids::random_generator gen; spider::TaskContext context = spider::core::TaskContextImpl::create_task_context( gen(), std::move(data_storage), - std::move(metadata_storage) + std::move(metadata_storage), + std::move(storage_factory) ); spider::core::FunctionManager const& manager = spider::core::FunctionManager::get_instance(); @@ -142,19 +157,19 @@ TEMPLATE_LIST_TEST_CASE( TEMPLATE_LIST_TEST_CASE( "Register and run function with data inputs", "[core][storage]", - spider::test::StorageTypeList + spider::test::StorageFactoryTypeList ) { - auto [unique_metadata_storage, unique_data_storage] = spider::test:: - create_storage, std::tuple_element_t<1, TestType>>(); - + std::shared_ptr const storage_factory + = spider::test::create_storage_factory(); std::shared_ptr const metadata_storage - = std::move(unique_metadata_storage); - std::shared_ptr const data_storage = std::move(unique_data_storage); + = storage_factory->provide_metadata_storage(); + std::shared_ptr const data_storage + = storage_factory->provide_data_storage(); - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_storage->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + auto conn = std::move(std::get>(conn_result)); msgpack::sbuffer buffer; msgpack::pack(buffer, 3); @@ -162,13 +177,14 @@ TEMPLATE_LIST_TEST_CASE( boost::uuids::random_generator gen; boost::uuids::uuid const driver_id = gen(); spider::core::Driver const driver{driver_id}; - REQUIRE(metadata_storage->add_driver(conn, driver).success()); - REQUIRE(data_storage->add_driver_data(conn, driver_id, data).success()); + REQUIRE(metadata_storage->add_driver(*conn, driver).success()); + REQUIRE(data_storage->add_driver_data(*conn, driver_id, data).success()); spider::TaskContext context = spider::core::TaskContextImpl::create_task_context( gen(), data_storage, - metadata_storage + metadata_storage, + storage_factory ); spider::core::FunctionManager const& manager = spider::core::FunctionManager::get_instance(); @@ -179,7 +195,7 @@ TEMPLATE_LIST_TEST_CASE( msgpack::sbuffer const result = (*function)(context, args_buffers); REQUIRE(3 == spider::core::response_get_result(result).value_or(0)); - REQUIRE(data_storage->remove_data(conn, data.get_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 dba48b7dd..e48b0c4bc 100644 --- a/tests/worker/test-TaskExecutor.cpp +++ b/tests/worker/test-TaskExecutor.cpp @@ -3,6 +3,7 @@ #include #include #include +#include #include #include @@ -22,7 +23,8 @@ #include "../../src/spider/io/MsgPack.hpp" // IWYU pragma: keep #include "../../src/spider/storage/DataStorage.hpp" #include "../../src/spider/storage/MetadataStorage.hpp" -#include "../../src/spider/storage/mysql/MySqlConnection.hpp" +#include "../../src/spider/storage/StorageConnection.hpp" +#include "../../src/spider/storage/StorageFactory.hpp" #include "../../src/spider/worker/FunctionManager.hpp" #include "../../src/spider/worker/TaskExecutor.hpp" #include "../storage/StorageTestHelper.hpp" @@ -58,7 +60,11 @@ auto get_libraries() -> std::vector { return {lib_path.string()}; } -TEST_CASE("Task execute success", "[worker][storage]") { +TEMPLATE_LIST_TEST_CASE( + "Task execute success", + "[worker][storage]", + spider::test::StorageFactoryTypeList +) { absl::flat_hash_map< boost::process::v2::environment::key, boost::process::v2::environment::value> const environment_variable @@ -72,7 +78,7 @@ TEST_CASE("Task execute success", "[worker][storage]") { context, "sum_test", gen(), - spider::test::cStorageUrl, + spider::test::get_storage_url(), get_libraries(), environment_variable, 2, @@ -86,7 +92,11 @@ TEST_CASE("Task execute success", "[worker][storage]") { REQUIRE(5 == result_option.value_or(0)); } -TEST_CASE("Task execute wrong number of arguments", "[worker][storage]") { +TEMPLATE_LIST_TEST_CASE( + "Task execute wrong number of arguments", + "[worker][storage]", + spider::test::StorageFactoryTypeList +) { absl::flat_hash_map< boost::process::v2::environment::key, boost::process::v2::environment::value> const environment_variable @@ -100,7 +110,7 @@ TEST_CASE("Task execute wrong number of arguments", "[worker][storage]") { context, "sum_test", gen(), - spider::test::cStorageUrl, + spider::test::get_storage_url(), get_libraries(), environment_variable, 2 @@ -112,7 +122,11 @@ TEST_CASE("Task execute wrong number of arguments", "[worker][storage]") { REQUIRE(spider::core::FunctionInvokeError::WrongNumberOfArguments == std::get<0>(error)); } -TEST_CASE("Task execute fail", "[worker][storage]") { +TEMPLATE_LIST_TEST_CASE( + "Task execute fail", + "[worker][storage]", + spider::test::StorageFactoryTypeList +) { absl::flat_hash_map< boost::process::v2::environment::key, boost::process::v2::environment::value> const environment_variable @@ -126,7 +140,7 @@ TEST_CASE("Task execute fail", "[worker][storage]") { context, "error_test", gen(), - spider::test::cStorageUrl, + spider::test::get_storage_url(), get_libraries(), environment_variable, 2 @@ -141,18 +155,19 @@ TEST_CASE("Task execute fail", "[worker][storage]") { TEMPLATE_LIST_TEST_CASE( "Task execute data argument", "[worker][storage]", - spider::test::StorageTypeList + spider::test::StorageFactoryTypeList ) { - auto [unique_metadata_storage, unique_data_storage] = spider::test:: - create_storage, std::tuple_element_t<1, TestType>>(); + std::shared_ptr const storage_factory + = spider::test::create_storage_factory(); std::shared_ptr const metadata_storage - = std::move(unique_metadata_storage); - std::shared_ptr const data_storage = std::move(unique_data_storage); + = storage_factory->provide_metadata_storage(); + std::shared_ptr const data_storage + = storage_factory->provide_data_storage(); - std::variant conn_result - = spider::core::MySqlConnection::create(metadata_storage->get_url()); - REQUIRE(std::holds_alternative(conn_result)); - auto& conn = std::get(conn_result); + std::variant, spider::core::StorageErr> + conn_result = storage_factory->provide_storage_connection(); + REQUIRE(std::holds_alternative>(conn_result)); + auto conn = std::move(std::get>(conn_result)); // Create driver and data msgpack::sbuffer buffer; @@ -161,8 +176,8 @@ TEMPLATE_LIST_TEST_CASE( boost::uuids::random_generator gen; boost::uuids::uuid const driver_id = gen(); spider::core::Driver const driver{driver_id}; - REQUIRE(metadata_storage->add_driver(conn, driver).success()); - REQUIRE(data_storage->add_driver_data(conn, driver_id, data).success()); + REQUIRE(metadata_storage->add_driver(*conn, driver).success()); + REQUIRE(data_storage->add_driver_data(*conn, driver_id, data).success()); absl::flat_hash_map< boost::process::v2::environment::key, @@ -175,7 +190,7 @@ TEMPLATE_LIST_TEST_CASE( context, "data_test", gen(), - spider::test::cStorageUrl, + spider::test::get_storage_url(), get_libraries(), environment_variable, data.get_id() @@ -190,7 +205,7 @@ TEMPLATE_LIST_TEST_CASE( } // Clean up - REQUIRE(data_storage->remove_data(conn, data.get_id()).success()); + REQUIRE(data_storage->remove_data(*conn, data.get_id()).success()); } } // namespace diff --git a/tools/scripts/storage/init_db.sql b/tools/scripts/storage/init_db.sql new file mode 100644 index 000000000..f002bb01c --- /dev/null +++ b/tools/scripts/storage/init_db.sql @@ -0,0 +1,149 @@ +CREATE TABLE IF NOT EXISTS `drivers` +( + `id` BINARY(16) NOT NULL, + `heartbeat` TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP, + PRIMARY KEY (`id`) +); +CREATE TABLE IF NOT EXISTS `schedulers` +( + `id` BINARY(16) NOT NULL, + `address` VARCHAR(40) NOT NULL, + `port` INT UNSIGNED NOT NULL, + `state` ENUM ('normal', 'recovery', 'gc') NOT NULL, + CONSTRAINT `scheduler_driver_id` FOREIGN KEY (`id`) REFERENCES `drivers` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE, + PRIMARY KEY (`id`) +); +CREATE TABLE IF NOT EXISTS jobs +( + `id` BINARY(16) NOT NULL, + `client_id` BINARY(16) NOT NULL, + `creation_time` TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP, + KEY (`client_id`) USING BTREE, + INDEX (`creation_time`), + PRIMARY KEY (`id`) +); +CREATE TABLE IF NOT EXISTS tasks +( + `id` BINARY(16) NOT NULL, + `job_id` BINARY(16) NOT NULL, + `func_name` VARCHAR(64) NOT NULL, + `state` ENUM ('pending', 'ready', 'running', 'success', 'cancel', 'fail') NOT NULL, + `timeout` FLOAT, + `max_retry` INT UNSIGNED DEFAULT 0, + `retry` INT UNSIGNED DEFAULT 0, + `instance_id` BINARY(16), + CONSTRAINT `task_job_id` FOREIGN KEY (`job_id`) REFERENCES `jobs` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE, + PRIMARY KEY (`id`) +); +CREATE TABLE IF NOT EXISTS input_tasks +( + `job_id` BINARY(16) NOT NULL, + `task_id` BINARY(16) NOT NULL, + `position` INT UNSIGNED NOT NULL, + CONSTRAINT `input_task_job_id` FOREIGN KEY (`job_id`) REFERENCES `jobs` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE, + CONSTRAINT `input_task_task_id` FOREIGN KEY (`task_id`) REFERENCES `tasks` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE, + INDEX (`job_id`, `position`), + PRIMARY KEY (`task_id`) +); +CREATE TABLE IF NOT EXISTS output_tasks +( + `job_id` BINARY(16) NOT NULL, + `task_id` BINARY(16) NOT NULL, + `position` INT UNSIGNED NOT NULL, + CONSTRAINT `output_task_job_id` FOREIGN KEY (`job_id`) REFERENCES `jobs` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE, + CONSTRAINT `output_task_task_id` FOREIGN KEY (`task_id`) REFERENCES `tasks` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE, + INDEX (`job_id`, `position`), + PRIMARY KEY (`task_id`) +); +CREATE TABLE IF NOT EXISTS `data` +( + `id` BINARY(16) NOT NULL, + `value` VARBINARY(256) NOT NULL, + `hard_locality` BOOL DEFAULT FALSE, + `persisted` BOOL DEFAULT FALSE, + PRIMARY KEY (`id`) +); +CREATE TABLE IF NOT EXISTS `task_outputs` +( + `task_id` BINARY(16) NOT NULL, + `position` INT UNSIGNED NOT NULL, + `type` VARCHAR(64) NOT NULL, + `value` VARBINARY(64), + `data_id` BINARY(16), + CONSTRAINT `output_task_id` FOREIGN KEY (`task_id`) REFERENCES `tasks` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE, + CONSTRAINT `output_data_id` FOREIGN KEY (`data_id`) REFERENCES `data` (`id`) ON UPDATE NO ACTION ON DELETE NO ACTION, + PRIMARY KEY (`task_id`, `position`) +); +CREATE TABLE IF NOT EXISTS `task_inputs` +( + `task_id` BINARY(16) NOT NULL, + `position` INT UNSIGNED NOT NULL, + `type` VARCHAR(64) NOT NULL, + `output_task_id` BINARY(16), + `output_task_position` INT UNSIGNED, + `value` VARBINARY(64), -- Use VARBINARY for all types of values + `data_id` BINARY(16), + CONSTRAINT `input_task_id` FOREIGN KEY (`task_id`) REFERENCES `tasks` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE, + CONSTRAINT `input_task_output_match` FOREIGN KEY (`output_task_id`, `output_task_position`) REFERENCES task_outputs (`task_id`, `position`) ON UPDATE NO ACTION ON DELETE SET NULL, + CONSTRAINT `input_data_id` FOREIGN KEY (`data_id`) REFERENCES `data` (`id`) ON UPDATE NO ACTION ON DELETE NO ACTION, + PRIMARY KEY (`task_id`, `position`) +); + +CREATE TABLE IF NOT EXISTS `task_dependencies` +( + `parent` BINARY(16) NOT NULL, + `child` BINARY(16) NOT NULL, + KEY (`parent`) USING BTREE, + KEY (`child`) USING BTREE, + CONSTRAINT `task_dep_parent` FOREIGN KEY (`parent`) REFERENCES `tasks` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE, + CONSTRAINT `task_dep_child` FOREIGN KEY (`child`) REFERENCES `tasks` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE +); +CREATE TABLE IF NOT EXISTS `task_instances` +( + `id` BINARY(16) NOT NULL, + `task_id` BINARY(16) NOT NULL, + `start_time` TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP ON UPDATE CURRENT_TIMESTAMP, + CONSTRAINT `instance_task_id` FOREIGN KEY (`task_id`) REFERENCES `tasks` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE, + PRIMARY KEY (`id`) +); + +CREATE TABLE IF NOT EXISTS `data_locality` +( + `id` BINARY(16) NOT NULL, + `address` VARCHAR(40) NOT NULL, + KEY (`id`) USING BTREE, + CONSTRAINT `locality_data_id` FOREIGN KEY (`id`) REFERENCES `data` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE +); +CREATE TABLE IF NOT EXISTS `data_ref_driver` +( + `id` BINARY(16) NOT NULL, + `driver_id` BINARY(16) NOT NULL, + KEY (`id`) USING BTREE, + KEY (`driver_id`) USING BTREE, + CONSTRAINT `data_driver_ref_id` FOREIGN KEY (`id`) REFERENCES `data` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE, + CONSTRAINT `data_ref_driver_id` FOREIGN KEY (`driver_id`) REFERENCES `drivers` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE +); +CREATE TABLE IF NOT EXISTS `data_ref_task` +( + `id` BINARY(16) NOT NULL, + `task_id` BINARY(16) NOT NULL, + KEY (`id`) USING BTREE, + KEY (`task_id`) USING BTREE, + CONSTRAINT `data_task_ref_id` FOREIGN KEY (`id`) REFERENCES `data` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE, + CONSTRAINT `data_ref_task_id` FOREIGN KEY (`task_id`) REFERENCES `tasks` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE +); +CREATE TABLE IF NOT EXISTS `client_kv_data` +( + `kv_key` VARCHAR(64) NOT NULL, + `value` VARBINARY(128) NOT NULL, + `client_id` BINARY(16) NOT NULL, + PRIMARY KEY (`client_id`, `kv_key`) +); +CREATE TABLE IF NOT EXISTS `task_kv_data` +( + `kv_key` VARCHAR(64) NOT NULL, + `value` VARBINARY(128) NOT NULL, + `task_id` BINARY(16) NOT NULL, + PRIMARY KEY (`task_id`, `kv_key`), + CONSTRAINT `kv_data_task_id` FOREIGN KEY (`task_id`) REFERENCES `tasks` (`id`) ON UPDATE NO ACTION ON DELETE CASCADE +); \ No newline at end of file