diff --git a/src/spider/CMakeLists.txt b/src/spider/CMakeLists.txt index 61e735a08..11f39e433 100644 --- a/src/spider/CMakeLists.txt +++ b/src/spider/CMakeLists.txt @@ -1,7 +1,7 @@ # set variable as CACHE INTERNAL to access it from other scope set(SPIDER_CORE_SOURCES - storage/MySqlConnection.cpp - storage/MySqlStorage.cpp + storage/mysql/MySqlConnection.cpp + storage/mysql/MySqlStorage.cpp worker/FunctionManager.cpp worker/FunctionNameManager.cpp io/msgpack_message.cpp @@ -25,8 +25,11 @@ set(SPIDER_CORE_HEADERS storage/MetadataStorage.hpp storage/DataStorage.hpp storage/StorageConnection.hpp - storage/MySqlConnection.hpp - storage/MySqlStorage.hpp + storage/mysql/mysql_stmt.hpp + storage/mysql/MySqlConnection.hpp + storage/mysql/MySqlStorage.hpp + storage/mysql/MySqlJobSubmissionBatch.hpp + storage/JobSubmissionBatch.hpp worker/FunctionManager.hpp worker/FunctionNameManager.hpp CACHE INTERNAL diff --git a/src/spider/client/Data.hpp b/src/spider/client/Data.hpp index 5ec729168..07cd1c72f 100644 --- a/src/spider/client/Data.hpp +++ b/src/spider/client/Data.hpp @@ -15,7 +15,7 @@ #include "../io/MsgPack.hpp" // IWYU pragma: keep #include "../io/Serializer.hpp" #include "../storage/DataStorage.hpp" -#include "../storage/MySqlConnection.hpp" +#include "../storage/mysql/MySqlConnection.hpp" #include "Exception.hpp" namespace spider { diff --git a/src/spider/client/Driver.cpp b/src/spider/client/Driver.cpp index ea59e08ca..b4d76b781 100644 --- a/src/spider/client/Driver.cpp +++ b/src/spider/client/Driver.cpp @@ -6,6 +6,7 @@ #include #include #include +#include #include #include @@ -15,8 +16,8 @@ #include "../core/Error.hpp" #include "../core/KeyValueData.hpp" #include "../io/BoostAsio.hpp" // IWYU pragma: keep -#include "../storage/MySqlConnection.hpp" -#include "../storage/MySqlStorage.hpp" +#include "../storage/mysql/MySqlConnection.hpp" +#include "../storage/mysql/MySqlStorage.hpp" #include "Exception.hpp" namespace spider { @@ -33,9 +34,11 @@ Driver::Driver(std::string const& storage_url) { if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - auto& conn = std::get(conn_result); + m_conn = std::make_shared( + std::get(std::move(conn_result)) + ); - core::StorageErr const err = m_metadata_storage->add_driver(conn, core::Driver{m_id}); + core::StorageErr const err = m_metadata_storage->add_driver(*m_conn, core::Driver{m_id}); if (!err.success()) { if (core::StorageErrType::DuplicateKeyErr == err.type) { throw DriverIdInUseException(m_id); @@ -71,9 +74,11 @@ Driver::Driver(std::string const& storage_url, boost::uuids::uuid const id) : m_ if (std::holds_alternative(conn_result)) { throw ConnectionException(std::get(conn_result).description); } - auto& conn = std::get(conn_result); + m_conn = std::make_shared( + std::get(std::move(conn_result)) + ); - core::StorageErr const err = m_metadata_storage->add_driver(conn, core::Driver{m_id}); + core::StorageErr const err = m_metadata_storage->add_driver(*m_conn, core::Driver{m_id}); if (!err.success()) { if (core::StorageErrType::DuplicateKeyErr == err.type) { throw DriverIdInUseException(m_id); @@ -104,29 +109,15 @@ Driver::Driver(std::string const& storage_url, boost::uuids::uuid const id) : m_ auto Driver::kv_store_insert(std::string const& key, std::string const& value) -> void { core::KeyValueData const kv_data{key, value, m_id}; - std::variant conn_result - = core::MySqlConnection::create(m_data_storage->get_url()); - if (std::holds_alternative(conn_result)) { - throw ConnectionException(std::get(conn_result).description); - } - auto& conn = std::get(conn_result); - - core::StorageErr const err = m_data_storage->add_client_kv_data(conn, kv_data); + core::StorageErr const err = m_data_storage->add_client_kv_data(*m_conn, kv_data); if (!err.success()) { throw ConnectionException(err.description); } } auto Driver::kv_store_get(std::string const& key) -> std::optional { - std::variant conn_result - = core::MySqlConnection::create(m_data_storage->get_url()); - if (std::holds_alternative(conn_result)) { - throw ConnectionException(std::get(conn_result).description); - } - auto& conn = std::get(conn_result); - std::string value; - core::StorageErr const err = m_data_storage->get_client_kv_data(conn, m_id, key, &value); + core::StorageErr const err = m_data_storage->get_client_kv_data(*m_conn, m_id, key, &value); if (!err.success()) { if (core::StorageErrType::KeyNotFoundErr == err.type) { return std::nullopt; diff --git a/src/spider/client/Driver.hpp b/src/spider/client/Driver.hpp index 84d1ffb62..ff5920a6b 100644 --- a/src/spider/client/Driver.hpp +++ b/src/spider/client/Driver.hpp @@ -8,7 +8,6 @@ #include #include #include -#include #include #include @@ -18,7 +17,10 @@ #include "../core/Error.hpp" #include "../core/TaskGraphImpl.hpp" #include "../io/Serializer.hpp" -#include "../storage/MySqlConnection.hpp" +#include "../storage/JobSubmissionBatch.hpp" +#include "../storage/mysql/MySqlConnection.hpp" +#include "../storage/mysql/MySqlJobSubmissionBatch.hpp" +#include "../storage/StorageConnection.hpp" #include "../worker/FunctionManager.hpp" #include "../worker/FunctionNameManager.hpp" #include "Data.hpp" @@ -133,6 +135,40 @@ class Driver { return TaskGraphType{std::move(graph)}; } + /** + * Begins a batch of `start` calls. This allows the driver to submit multiple jobs in a single + * batch, which can be more efficient than submitting jobs individually. + * + * Needs to be paired with `end_batch_start`. + * + * If a batch has already been started, this method is a no-op. + */ + auto begin_batch_start() -> void { + if (nullptr != m_batch) { + return; + } + m_batch = std::make_shared( + // NOLINTNEXTLINE(cppcoreguidelines-pro-type-static-cast-downcast) + static_cast(*m_conn) + ); + } + + /** + * Ends a batch of `start` calls. This submits all jobs in the batch to Spider. + * + * @throw spider::ConnectionException + */ + auto end_batch_start() -> void { + if (nullptr == m_batch) { + return; + } + core::StorageErr const err = m_batch->submit_batch(*m_conn); + m_batch = nullptr; + if (!err.success()) { + throw ConnectionException(fmt::format("Failed to start job: {}", err.description)); + } + } + /** * Starts running a task with the given inputs on Spider. * @@ -175,18 +211,20 @@ class Driver { graph.add_task(new_task); graph.add_input_task(new_task.get_id()); graph.add_output_task(new_task.get_id()); - std::variant conn_result - = core::MySqlConnection::create(m_metadata_storage->get_url()); - if (std::holds_alternative(conn_result)) { - throw ConnectionException(std::get(conn_result).description); - } - auto& conn = std::get(conn_result); - core::StorageErr err = m_metadata_storage->add_job(conn, job_id, m_id, graph); - if (!err.success()) { - throw ConnectionException(fmt::format("Failed to start job: {}", err.description)); + if (nullptr != m_batch) { + core::StorageErr const err + = m_metadata_storage->add_job_batch(*m_conn, *m_batch, job_id, m_id, graph); + if (!err.success()) { + throw ConnectionException(fmt::format("Failed to start job: {}", err.description)); + } + } else { + core::StorageErr const err = m_metadata_storage->add_job(*m_conn, job_id, m_id, graph); + if (!err.success()) { + throw ConnectionException(fmt::format("Failed to start job: {}", err.description)); + } } - return Job{job_id, m_metadata_storage, m_data_storage}; + return Job{job_id, m_metadata_storage, m_data_storage, m_conn}; } /** @@ -224,19 +262,13 @@ class Driver { graph.m_impl->reset_ids(); boost::uuids::random_generator gen; boost::uuids::uuid const job_id = gen(); - std::variant conn_result - = core::MySqlConnection::create(m_metadata_storage->get_url()); - if (std::holds_alternative(conn_result)) { - throw ConnectionException(std::get(conn_result).description); - } - auto& conn = std::get(conn_result); core::StorageErr const err - = m_metadata_storage->add_job(conn, job_id, m_id, graph.m_impl->get_graph()); + = m_metadata_storage->add_job(*m_conn, job_id, m_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_storage, m_data_storage}; + return Job{job_id, m_metadata_storage, m_data_storage, m_conn}; } /** @@ -249,14 +281,8 @@ class Driver { */ auto get_jobs() -> std::vector { std::vector job_ids; - std::variant conn_result - = core::MySqlConnection::create(m_metadata_storage->get_url()); - if (std::holds_alternative(conn_result)) { - throw ConnectionException(std::get(conn_result).description); - } - auto& conn = std::get(conn_result); core::StorageErr const err - = m_metadata_storage->get_jobs_by_client_id(conn, m_id, &job_ids); + = m_metadata_storage->get_jobs_by_client_id(*m_conn, m_id, &job_ids); if (!err.success()) { throw ConnectionException("Failed to get jobs."); } @@ -267,6 +293,8 @@ class Driver { boost::uuids::uuid m_id; std::shared_ptr m_metadata_storage; std::shared_ptr m_data_storage; + std::shared_ptr m_conn; + std::shared_ptr m_batch{nullptr}; std::jthread m_heartbeat_thread; }; } // namespace spider diff --git a/src/spider/client/Job.hpp b/src/spider/client/Job.hpp index 8b0ff626f..07be6a6c6 100644 --- a/src/spider/client/Job.hpp +++ b/src/spider/client/Job.hpp @@ -21,7 +21,8 @@ #include "../core/JobMetadata.hpp" #include "../io/MsgPack.hpp" // IWYU pragma: keep #include "../storage/MetadataStorage.hpp" -#include "../storage/MySqlConnection.hpp" +#include "../storage/mysql/MySqlConnection.hpp" +#include "../storage/StorageConnection.hpp" #include "Data.hpp" #include "Exception.hpp" #include "task.hpp" @@ -63,29 +64,16 @@ class Job { * @throw spider::ConnectionException */ auto wait_complete() -> void { - std::variant conn_result - = core::MySqlConnection::create(m_data_storage->get_url()); - if (std::holds_alternative(conn_result)) { - throw ConnectionException(std::get(conn_result).description); - } - auto& conn = std::get(conn_result); - - bool complete = false; - core::StorageErr err = m_metadata_storage->get_job_complete(conn, m_id, &complete); - if (!err.success()) { - throw ConnectionException{ - fmt::format("Failed to get job completion status: {}", err.description) - }; - } - while (!complete) { - constexpr int cSleepMs = 10; - std::this_thread::sleep_for(std::chrono::milliseconds(cSleepMs)); - err = m_metadata_storage->get_job_complete(conn, m_id, &complete); - if (!err.success()) { - throw ConnectionException{ - fmt::format("Failed to get job completion status: {}", err.description) - }; + if (nullptr == m_conn) { + std::variant conn_result + = core::MySqlConnection::create(m_data_storage->get_url()); + if (std::holds_alternative(conn_result)) { + throw ConnectionException(std::get(conn_result).description); } + auto& conn = std::get(conn_result); + wait_complete_conn(conn); + } else { + wait_complete_conn(*m_conn); } } @@ -101,15 +89,21 @@ class Job { * @throw spider::ConnectionException */ auto get_status() -> JobStatus { - std::variant conn_result - = core::MySqlConnection::create(m_data_storage->get_url()); - if (std::holds_alternative(conn_result)) { - throw ConnectionException(std::get(conn_result).description); - } - auto& conn = std::get(conn_result); - core::JobStatus status = core::JobStatus::Running; - core::StorageErr const err = m_metadata_storage->get_job_status(conn, m_id, &status); + core::StorageErr err; + + if (nullptr == m_conn) { + std::variant conn_result + = core::MySqlConnection::create(m_data_storage->get_url()); + if (std::holds_alternative(conn_result)) { + throw ConnectionException(std::get(conn_result).description); + } + auto& conn = std::get(conn_result); + + err = m_metadata_storage->get_job_status(conn, m_id, &status); + } else { + err = m_metadata_storage->get_job_status(*m_conn, m_id, &status); + } if (!err.success()) { throw ConnectionException{fmt::format("Failed to get job status: {}", err.description)}; } @@ -128,7 +122,6 @@ class Job { }; } - // NOLINTBEGIN(readability-function-cognitive-complexity) /** * NOTE: It is undefined behavior to call this method for a job that is not in the `Succeeded` * state. @@ -137,13 +130,70 @@ class Job { * @throw spider::ConnectionException */ auto get_result() -> ReturnType { - std::variant conn_result - = core::MySqlConnection::create(m_data_storage->get_url()); - if (std::holds_alternative(conn_result)) { - throw ConnectionException(std::get(conn_result).description); + if (nullptr == m_conn) { + std::variant conn_result + = core::MySqlConnection::create(m_data_storage->get_url()); + 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::get(conn_result); + return get_result_conn(*m_conn); + } + + /** + * NOTE: It is undefined behavior to call this method for a job that is not in the `Failed` + * state. + * + * @return A pair: + * - the name of the task function that failed. + * - the error message sent from the task through `TaskContext::abort` or from Spider. + * @throw spider::ConnectionException + */ + auto get_error() -> std::pair { + throw ConnectionException{"Not implemented"}; + } +private: + Job(boost::uuids::uuid id, + std::shared_ptr metadata_storage, + std::shared_ptr data_storage) + : m_id{id}, + m_metadata_storage{std::move(metadata_storage)}, + m_data_storage{std::move(data_storage)} {} + + Job(boost::uuids::uuid id, + std::shared_ptr metadata_storage, + std::shared_ptr data_storage, + std::shared_ptr conn) + : m_id{id}, + m_metadata_storage{std::move(metadata_storage)}, + m_data_storage{std::move(data_storage)}, + m_conn{std::move(conn)} {} + + auto wait_complete_conn(core::StorageConnection& conn) -> void { + bool complete = false; + core::StorageErr err = m_metadata_storage->get_job_complete(conn, m_id, &complete); + if (!err.success()) { + throw ConnectionException{ + fmt::format("Failed to get job completion status: {}", err.description) + }; + } + while (!complete) { + constexpr int cSleepMs = 10; + std::this_thread::sleep_for(std::chrono::milliseconds(cSleepMs)); + err = m_metadata_storage->get_job_complete(conn, m_id, &complete); + if (!err.success()) { + throw ConnectionException{ + fmt::format("Failed to get job completion status: {}", err.description) + }; + } + } + } + + // NOLINTBEGIN(readability-function-cognitive-complexity) + auto get_result_conn(core::StorageConnection& conn) -> ReturnType { std::vector output_task_ids; core::StorageErr err = m_metadata_storage->get_job_output_tasks(conn, m_id, &output_task_ids); @@ -277,30 +327,10 @@ class Job { // NOLINTEND(readability-function-cognitive-complexity) - /** - * NOTE: It is undefined behavior to call this method for a job that is not in the `Failed` - * state. - * - * @return A pair: - * - the name of the task function that failed. - * - the error message sent from the task through `TaskContext::abort` or from Spider. - * @throw spider::ConnectionException - */ - auto get_error() -> std::pair { - throw ConnectionException{"Not implemented"}; - } - -private: - Job(boost::uuids::uuid id, - std::shared_ptr metadata_storage, - std::shared_ptr data_storage) - : m_id{id}, - m_metadata_storage{std::move(metadata_storage)}, - m_data_storage{std::move(data_storage)} {} - boost::uuids::uuid m_id; std::shared_ptr m_metadata_storage; std::shared_ptr m_data_storage; + std::shared_ptr m_conn; friend class Driver; friend class TaskContext; diff --git a/src/spider/client/TaskContext.cpp b/src/spider/client/TaskContext.cpp index 8aa4d59b2..f34f4f37c 100644 --- a/src/spider/client/TaskContext.cpp +++ b/src/spider/client/TaskContext.cpp @@ -9,7 +9,7 @@ #include "../core/Error.hpp" #include "../core/KeyValueData.hpp" -#include "../storage/MySqlConnection.hpp" +#include "../storage/mysql/MySqlConnection.hpp" #include "Exception.hpp" namespace spider { diff --git a/src/spider/client/TaskContext.hpp b/src/spider/client/TaskContext.hpp index b370f6a85..ce80c2bec 100644 --- a/src/spider/client/TaskContext.hpp +++ b/src/spider/client/TaskContext.hpp @@ -19,7 +19,7 @@ #include "../core/TaskGraph.hpp" #include "../core/TaskGraphImpl.hpp" #include "../io/Serializer.hpp" -#include "../storage/MySqlConnection.hpp" +#include "../storage/mysql/MySqlConnection.hpp" #include "Data.hpp" #include "Exception.hpp" #include "Job.hpp" diff --git a/src/spider/scheduler/scheduler.cpp b/src/spider/scheduler/scheduler.cpp index 3ac8e85ea..25135074c 100644 --- a/src/spider/scheduler/scheduler.cpp +++ b/src/spider/scheduler/scheduler.cpp @@ -24,8 +24,8 @@ #include "../io/BoostAsio.hpp" // IWYU pragma: keep #include "../storage/DataStorage.hpp" #include "../storage/MetadataStorage.hpp" -#include "../storage/MySqlConnection.hpp" -#include "../storage/MySqlStorage.hpp" +#include "../storage/mysql/MySqlConnection.hpp" +#include "../storage/mysql/MySqlStorage.hpp" #include "../utils/StopToken.hpp" #include "FifoPolicy.hpp" #include "SchedulerPolicy.hpp" diff --git a/src/spider/storage/JobSubmissionBatch.hpp b/src/spider/storage/JobSubmissionBatch.hpp new file mode 100644 index 000000000..30852a947 --- /dev/null +++ b/src/spider/storage/JobSubmissionBatch.hpp @@ -0,0 +1,21 @@ +#ifndef SPIDER_STORAGE_JOBSUBMISSIONBATCH_HPP +#define SPIDER_STORAGE_JOBSUBMISSIONBATCH_HPP + +#include "../core/Error.hpp" +#include "StorageConnection.hpp" + +namespace spider::core { +class JobSubmissionBatch { +public: + virtual auto submit_batch(StorageConnection& conn) -> StorageErr = 0; + + JobSubmissionBatch() = default; + JobSubmissionBatch(JobSubmissionBatch const&) = delete; + auto operator=(JobSubmissionBatch const&) -> JobSubmissionBatch& = delete; + JobSubmissionBatch(JobSubmissionBatch&&) = delete; + auto operator=(JobSubmissionBatch&&) -> JobSubmissionBatch& = delete; + virtual ~JobSubmissionBatch() = default; +}; +} // namespace spider::core + +#endif diff --git a/src/spider/storage/MetadataStorage.hpp b/src/spider/storage/MetadataStorage.hpp index 47564f918..d1a1d169d 100644 --- a/src/spider/storage/MetadataStorage.hpp +++ b/src/spider/storage/MetadataStorage.hpp @@ -12,6 +12,7 @@ #include "../core/JobMetadata.hpp" #include "../core/Task.hpp" #include "../core/TaskGraph.hpp" +#include "JobSubmissionBatch.hpp" #include "StorageConnection.hpp" namespace spider::core { @@ -38,6 +39,13 @@ class MetadataStorage { boost::uuids::uuid client_id, TaskGraph const& task_graph ) -> StorageErr = 0; + virtual auto add_job_batch( + StorageConnection& conn, + JobSubmissionBatch& batch, + boost::uuids::uuid job_id, + boost::uuids::uuid client_id, + TaskGraph const& task_graph + ) -> StorageErr = 0; virtual auto get_job_metadata(StorageConnection& conn, boost::uuids::uuid id, JobMetadata* job) -> StorageErr = 0; virtual auto get_job_complete(StorageConnection& conn, boost::uuids::uuid id, bool* complete) diff --git a/src/spider/storage/MySqlConnection.cpp b/src/spider/storage/mysql/MySqlConnection.cpp similarity index 66% rename from src/spider/storage/MySqlConnection.cpp rename to src/spider/storage/mysql/MySqlConnection.cpp index ab751c58a..6cc56c1f0 100644 --- a/src/spider/storage/MySqlConnection.cpp +++ b/src/spider/storage/mysql/MySqlConnection.cpp @@ -7,35 +7,25 @@ #include #include -#include +#include #include -#include +#include #include -#include "../core/Error.hpp" +#include "../../core/Error.hpp" namespace spider::core { auto MySqlConnection::create(std::string const& url) -> std::variant { - // Parse jdbc url + // Validate jdbc url std::regex const url_regex(R"(jdbc:mariadb://[^?]+(\?user=([^&]*)(&password=([^&]*))?)?)"); std::smatch match; if (false == std::regex_match(url, match, url_regex)) { return StorageErr{StorageErrType::OtherErr, "Invalid url"}; } - bool const credential = match[2].matched && match[4].matched; - std::unique_ptr conn; try { - sql::Driver* driver = sql::mariadb::get_driver_instance(); - if (credential) { - conn = std::unique_ptr( - driver->connect(sql::SQLString(url), match[2].str(), match[4].str()) - ); - } else { - conn = std::unique_ptr( - driver->connect(sql::SQLString(url), sql::Properties{}) - ); - } + sql::Properties const properties{{{"useBulkStmts", "true"}}}; + std::unique_ptr conn{sql::DriverManager::getConnection(url, properties)}; conn->setAutoCommit(false); return MySqlConnection{std::move(conn)}; } catch (sql::SQLException& e) { diff --git a/src/spider/storage/MySqlConnection.hpp b/src/spider/storage/mysql/MySqlConnection.hpp similarity index 94% rename from src/spider/storage/MySqlConnection.hpp rename to src/spider/storage/mysql/MySqlConnection.hpp index 1bfc65001..5c9cb2cc7 100644 --- a/src/spider/storage/MySqlConnection.hpp +++ b/src/spider/storage/mysql/MySqlConnection.hpp @@ -8,8 +8,8 @@ #include -#include "../core/Error.hpp" -#include "StorageConnection.hpp" +#include "../../core/Error.hpp" +#include "../StorageConnection.hpp" namespace spider::core { diff --git a/src/spider/storage/mysql/MySqlJobSubmissionBatch.hpp b/src/spider/storage/mysql/MySqlJobSubmissionBatch.hpp new file mode 100644 index 000000000..327543fa2 --- /dev/null +++ b/src/spider/storage/mysql/MySqlJobSubmissionBatch.hpp @@ -0,0 +1,95 @@ +#ifndef SPIDER_STORAGE_MYSQLJOBSUBMISSIONBATCH_HPP +#define SPIDER_STORAGE_MYSQLJOBSUBMISSIONBATCH_HPP + +#include + +#include +#include +#include + +#include "../../core/Error.hpp" +#include "../JobSubmissionBatch.hpp" +#include "../StorageConnection.hpp" +#include "mysql_stmt.hpp" +#include "MySqlConnection.hpp" + +namespace spider::core { +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{}; + } + + auto get_job_stmt() -> sql::PreparedStatement& { return *m_job_stmt; } + + auto get_task_stmt() -> sql::PreparedStatement& { return *m_task_stmt; } + + auto get_task_input_output_stmt() -> sql::PreparedStatement& { + return *m_task_input_output_stmt; + } + + auto get_task_input_value_stmt() -> sql::PreparedStatement& { return *m_task_input_value_stmt; } + + auto get_task_input_data_stmt() -> sql::PreparedStatement& { return *m_task_input_data_stmt; } + + auto get_task_output_stmt() -> sql::PreparedStatement& { return *m_task_output_stmt; } + + auto get_task_dependency_stmt() -> sql::PreparedStatement& { return *m_task_dependency_stmt; } + + auto get_input_task_stmt() -> sql::PreparedStatement& { return *m_input_task_stmt; } + + auto get_output_task_stmt() -> sql::PreparedStatement& { return *m_output_task_stmt; } + +private: + std::unique_ptr m_job_stmt; + std::unique_ptr m_task_stmt; + std::unique_ptr m_task_input_output_stmt; + std::unique_ptr m_task_input_value_stmt; + std::unique_ptr m_task_input_data_stmt; + std::unique_ptr m_task_output_stmt; + std::unique_ptr m_task_dependency_stmt; + std::unique_ptr m_input_task_stmt; + std::unique_ptr m_output_task_stmt; +}; +} // namespace spider::core + +#endif diff --git a/src/spider/storage/MySqlStorage.cpp b/src/spider/storage/mysql/MySqlStorage.cpp similarity index 88% rename from src/spider/storage/MySqlStorage.cpp rename to src/spider/storage/mysql/MySqlStorage.cpp index 012eed730..06e4f7860 100644 --- a/src/spider/storage/MySqlStorage.cpp +++ b/src/spider/storage/mysql/MySqlStorage.cpp @@ -1,7 +1,6 @@ #include "MySqlStorage.hpp" #include -#include #include #include #include @@ -28,15 +27,18 @@ #include #include -#include "../core/Data.hpp" -#include "../core/Driver.hpp" -#include "../core/Error.hpp" -#include "../core/JobMetadata.hpp" -#include "../core/KeyValueData.hpp" -#include "../core/Task.hpp" -#include "../core/TaskGraph.hpp" +#include "../../core/Data.hpp" +#include "../../core/Driver.hpp" +#include "../../core/Error.hpp" +#include "../../core/JobMetadata.hpp" +#include "../../core/KeyValueData.hpp" +#include "../../core/Task.hpp" +#include "../../core/TaskGraph.hpp" +#include "../JobSubmissionBatch.hpp" +#include "../StorageConnection.hpp" +#include "mysql_stmt.hpp" #include "MySqlConnection.hpp" -#include "StorageConnection.hpp" +#include "MySqlJobSubmissionBatch.hpp" // mariadb-connector-cpp does not define SQL errcode. Just include some useful ones. enum MariadbErr : uint16_t { @@ -52,171 +54,6 @@ enum MariadbErr : uint16_t { namespace spider::core { namespace { -char const* const cCreateDriverTable = R"(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`) -))"; - -char const* const cCreateSchedulerTable = R"(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`) -))"; - -char const* const cCreateJobTable = R"(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`) -))"; - -char const* const cCreateTaskTable = R"(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`) -))"; - -char const* const cCreateInputTaskTable = R"(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`) -))"; - -char const* const cCreateOutputTaskTable = R"(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`) -))"; - -char const* const cCreateTaskInputTable = R"(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`) -))"; - -char const* const cCreateTaskOutputTable = R"(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`) -))"; - -char const* const cCreateTaskDependencyTable = R"(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 -))"; - -char const* const cCreateTaskInstanceTable = R"(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`) -))"; - -char const* const cCreateDataTable = R"(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`) -))"; - -char const* const cCreateDataLocalityTable = R"(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 -))"; - -char const* const cCreateDataRefDriverTable = R"(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 -))"; - -char const* const cCreateDataRefTaskTable = R"(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 -))"; - -char const* const cCreateClientKVDataTable = R"(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`) -))"; - -char const* const cCreateTaskKVDataTable = R"(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 -))"; - -std::array const cCreateStorage = { - cCreateDriverTable, // drivers table must be created before data_ref_driver - cCreateSchedulerTable, - cCreateJobTable, // jobs table must be created before task - cCreateTaskTable, // tasks table must be created before data_ref_task - cCreateDataTable, // data table must be created before task_outputs - cCreateDataLocalityTable, - cCreateDataRefDriverTable, - cCreateDataRefTaskTable, - cCreateClientKVDataTable, - cCreateTaskKVDataTable, - cCreateInputTaskTable, - cCreateOutputTaskTable, - cCreateTaskOutputTable, // task_outputs table must be created before task_inputs - cCreateTaskInputTable, - cCreateTaskDependencyTable, - cCreateTaskInstanceTable, -}; auto uuid_get_bytes(boost::uuids::uuid const& id) -> sql::bytes { // NOLINTBEGIN(cppcoreguidelines-pro-type-cstyle-cast) @@ -287,7 +124,7 @@ auto string_to_task_state(std::string const& state) -> spider::core::TaskState { // NOLINTBEGIN(cppcoreguidelines-pro-type-static-cast-downcast) auto MySqlMetadataStorage::initialize(StorageConnection& conn) -> StorageErr { try { - for (char const* create_table_str : cCreateStorage) { + for (std::string const& create_table_str : mysql::cCreateStorage) { std::unique_ptr statement( static_cast(conn)->createStatement() ); @@ -381,18 +218,25 @@ auto MySqlMetadataStorage::get_active_scheduler( return StorageErr{}; } -void MySqlMetadataStorage::add_task(MySqlConnection& conn, sql::bytes job_id, Task const& task) { +void MySqlMetadataStorage::add_task( + MySqlConnection& conn, + sql::bytes job_id, + Task const& task, + std::optional const& state +) { // Add task - std::unique_ptr task_statement( - conn->prepareStatement("INSERT INTO `tasks` (`id`, `job_id`, `func_name`, `state`, " - "`timeout`, `max_retry`) VALUES (?, ?, ?, ?, ?, ?)") - ); + std::unique_ptr task_statement(conn->prepareStatement(mysql::cInsertTask + )); sql::bytes task_id_bytes = uuid_get_bytes(task.get_id()); // NOLINTBEGIN(cppcoreguidelines-avoid-magic-numbers, readability-magic-numbers) task_statement->setBytes(1, &task_id_bytes); task_statement->setBytes(2, &job_id); task_statement->setString(3, task.get_function_name()); - task_statement->setString(4, task_state_to_string(task.get_state())); + if (state.has_value()) { + task_statement->setString(4, task_state_to_string(state.value())); + } else { + task_statement->setString(4, task_state_to_string(task.get_state())); + } task_statement->setFloat(5, task.get_timeout()); task_statement->setUInt(6, task.get_max_retries()); // NOLINTEND(cppcoreguidelines-avoid-magic-numbers, readability-magic-numbers) @@ -407,10 +251,9 @@ void MySqlMetadataStorage::add_task(MySqlConnection& conn, sql::bytes job_id, Ta std::optional const& value = input.get_value(); if (task_output.has_value()) { std::tuple const pair = task_output.value(); - std::unique_ptr input_statement(conn->prepareStatement( - "INSERT INTO `task_inputs` (`task_id`, `position`, `type`, `output_task_id`, " - "`output_task_position`) VALUES (?, ?, ?, ?, ?)" - )); + std::unique_ptr input_statement( + conn->prepareStatement(mysql::cInsertTaskInputOutput) + ); // NOLINTBEGIN(cppcoreguidelines-avoid-magic-numbers, readability-magic-numbers) input_statement->setBytes(1, &task_id_bytes); input_statement->setUInt(2, i); @@ -422,8 +265,7 @@ void MySqlMetadataStorage::add_task(MySqlConnection& conn, sql::bytes job_id, Ta input_statement->executeUpdate(); } else if (data_id.has_value()) { std::unique_ptr input_statement( - conn->prepareStatement("INSERT INTO `task_inputs` (`task_id`, `position`, " - "`type`, `data_id`) VALUES (?, ?, ?, ?)") + conn->prepareStatement(mysql::cInsertTaskInputData) ); input_statement->setBytes(1, &task_id_bytes); input_statement->setUInt(2, i); @@ -433,8 +275,7 @@ void MySqlMetadataStorage::add_task(MySqlConnection& conn, sql::bytes job_id, Ta input_statement->executeUpdate(); } else if (value.has_value()) { std::unique_ptr input_statement( - conn->prepareStatement("INSERT INTO `task_inputs` (`task_id`, `position`, " - "`type`, `value`) VALUES (?, ?, ?, ?)") + conn->prepareStatement(mysql::cInsertTaskInputValue) ); input_statement->setBytes(1, &task_id_bytes); input_statement->setUInt(2, i); @@ -447,9 +288,9 @@ void MySqlMetadataStorage::add_task(MySqlConnection& conn, sql::bytes job_id, Ta // Add task outputs for (std::uint64_t i = 0; i < task.get_num_outputs(); i++) { TaskOutput const output = task.get_output(i); - std::unique_ptr output_statement(conn->prepareStatement( - "INSERT INTO `task_outputs` (`task_id`, `position`, `type`) VALUES (?, ?, ?)" - )); + std::unique_ptr output_statement( + conn->prepareStatement(mysql::cInsertTaskOutput) + ); output_statement->setBytes(1, &task_id_bytes); output_statement->setUInt(2, i); output_statement->setString(3, output.get_type()); @@ -457,6 +298,77 @@ void MySqlMetadataStorage::add_task(MySqlConnection& conn, sql::bytes job_id, Ta } } +void MySqlMetadataStorage::add_task_batch( + MySqlJobSubmissionBatch& batch, + sql::bytes job_id, + Task const& task, + std::optional const& state +) { + // Add task + sql::PreparedStatement& task_statement = batch.get_task_stmt(); + sql::bytes task_id_bytes = uuid_get_bytes(task.get_id()); + // NOLINTBEGIN(cppcoreguidelines-avoid-magic-numbers, readability-magic-numbers) + task_statement.setBytes(1, &task_id_bytes); + task_statement.setBytes(2, &job_id); + task_statement.setString(3, task.get_function_name()); + if (state.has_value()) { + task_statement.setString(4, task_state_to_string(state.value())); + } else { + task_statement.setString(4, task_state_to_string(task.get_state())); + } + task_statement.setFloat(5, task.get_timeout()); + task_statement.setUInt(6, task.get_max_retries()); + // NOLINTEND(cppcoreguidelines-avoid-magic-numbers, readability-magic-numbers) + task_statement.addBatch(); + + // Add task inputs + for (std::uint64_t i = 0; i < task.get_num_inputs(); ++i) { + TaskInput const input = task.get_input(i); + std::optional> const task_output + = input.get_task_output(); + std::optional const data_id = input.get_data_id(); + std::optional const& value = input.get_value(); + if (task_output.has_value()) { + std::tuple const pair = task_output.value(); + sql::PreparedStatement& input_statement = batch.get_task_input_output_stmt(); + // NOLINTBEGIN(cppcoreguidelines-avoid-magic-numbers, readability-magic-numbers) + input_statement.setBytes(1, &task_id_bytes); + input_statement.setUInt(2, i); + input_statement.setString(3, input.get_type()); + sql::bytes task_output_id = uuid_get_bytes(std::get<0>(pair)); + input_statement.setBytes(4, &task_output_id); + input_statement.setUInt(5, std::get<1>(pair)); + // NOLINTEND(cppcoreguidelines-avoid-magic-numbers, readability-magic-numbers) + input_statement.addBatch(); + } else if (data_id.has_value()) { + sql::PreparedStatement& input_statement = batch.get_task_input_data_stmt(); + input_statement.setBytes(1, &task_id_bytes); + input_statement.setUInt(2, i); + input_statement.setString(3, input.get_type()); + sql::bytes data_id_bytes = uuid_get_bytes(data_id.value()); + input_statement.setBytes(4, &data_id_bytes); + input_statement.addBatch(); + } else if (value.has_value()) { + sql::PreparedStatement& input_statement = batch.get_task_input_value_stmt(); + input_statement.setBytes(1, &task_id_bytes); + input_statement.setUInt(2, i); + input_statement.setString(3, input.get_type()); + input_statement.setString(4, value.value()); + input_statement.addBatch(); + } + } + + // Add task outputs + for (std::uint64_t i = 0; i < task.get_num_outputs(); i++) { + TaskOutput const output = task.get_output(i); + sql::PreparedStatement& output_statement = batch.get_task_output_stmt(); + output_statement.setBytes(1, &task_id_bytes); + output_statement.setUInt(2, i); + output_statement.setString(3, output.get_type()); + output_statement.addBatch(); + } +} + // NOLINTBEGIN(readability-function-cognitive-complexity) auto MySqlMetadataStorage::add_job( StorageConnection& conn, @@ -469,9 +381,7 @@ auto MySqlMetadataStorage::add_job( sql::bytes client_id_bytes = uuid_get_bytes(client_id); { std::unique_ptr statement{ - static_cast(conn)->prepareStatement( - "INSERT INTO `jobs` (`id`, `client_id`) VALUES (?, ?)" - ) + static_cast(conn)->prepareStatement(mysql::cInsertJob) }; statement->setBytes(1, &job_id_bytes); statement->setBytes(2, &client_id_bytes); @@ -496,7 +406,7 @@ auto MySqlMetadataStorage::add_job( }; } Task const* task = task_option.value(); - add_task(static_cast(conn), job_id_bytes, *task); + add_task(static_cast(conn), job_id_bytes, *task, TaskState::Ready); for (boost::uuids::uuid const id : task_graph.get_child_tasks(task_id)) { std::vector const parents = task_graph.get_parent_tasks(id); if (std::ranges::all_of(parents, [&](boost::uuids::uuid const& parent) { @@ -519,7 +429,7 @@ auto MySqlMetadataStorage::add_job( return StorageErr{StorageErrType::KeyNotFoundErr, "Task graph inconsistent"}; } Task const* task = task_option.value(); - add_task(static_cast(conn), job_id_bytes, *task); + add_task(static_cast(conn), job_id_bytes, *task, std::nullopt); for (boost::uuids::uuid const id : task_graph.get_child_tasks(task_id)) { std::vector const parents = task_graph.get_parent_tasks(id); if (std::ranges::all_of(parents, [&](boost::uuids::uuid const& parent) { @@ -538,7 +448,7 @@ auto MySqlMetadataStorage::add_job( { std::unique_ptr dep_statement{ static_cast(conn)->prepareStatement( - "INSERT INTO `task_dependencies` (parent, child) VALUES (?, ?)" + mysql::cInsertTaskDependency ) }; sql::bytes parent_id_bytes = uuid_get_bytes(pair.first); @@ -551,10 +461,7 @@ auto MySqlMetadataStorage::add_job( // Add input tasks for (size_t i = 0; i < input_task_ids.size(); i++) { std::unique_ptr input_statement{ - static_cast(conn)->prepareStatement( - "INSERT INTO `input_tasks` (`job_id`, `task_id`, `position`) VALUES " - "(?, ?, ?)" - ) + static_cast(conn)->prepareStatement(mysql::cInsertInputTask) }; input_statement->setBytes(1, &job_id_bytes); sql::bytes task_id_bytes = uuid_get_bytes(input_task_ids[i]); @@ -566,10 +473,7 @@ auto MySqlMetadataStorage::add_job( std::vector const& output_task_ids = task_graph.get_output_tasks(); for (size_t i = 0; i < output_task_ids.size(); i++) { std::unique_ptr output_statement{ - static_cast(conn)->prepareStatement( - "INSERT INTO `output_tasks` (`job_id`, `task_id`, `position`) VALUES " - "(?, ?, ?)" - ) + static_cast(conn)->prepareStatement(mysql::cInsertOutputTask) }; output_statement->setBytes(1, &job_id_bytes); sql::bytes task_id_bytes = uuid_get_bytes(output_task_ids[i]); @@ -578,16 +482,132 @@ auto MySqlMetadataStorage::add_job( output_statement->executeUpdate(); } - // Mark head tasks as ready - for (boost::uuids::uuid const& task_id : task_graph.get_input_tasks()) { - std::unique_ptr statement( - static_cast(conn)->prepareStatement( - "UPDATE `tasks` SET `state` = 'ready' WHERE `id` = ?" - ) + } catch (sql::SQLException& e) { + static_cast(conn)->rollback(); + if (e.getErrorCode() == ErDupKey || e.getErrorCode() == ErDupEntry) { + return StorageErr{StorageErrType::DuplicateKeyErr, e.what()}; + } + return StorageErr{StorageErrType::OtherErr, e.what()}; + } + static_cast(conn)->commit(); + return StorageErr{}; +} + +auto MySqlMetadataStorage::add_job_batch( + StorageConnection& conn, + JobSubmissionBatch& batch, + boost::uuids::uuid job_id, + boost::uuids::uuid client_id, + TaskGraph const& task_graph +) -> StorageErr { + try { + sql::bytes job_id_bytes = uuid_get_bytes(job_id); + sql::bytes client_id_bytes = uuid_get_bytes(client_id); + { + sql::PreparedStatement& statement + = static_cast(batch).get_job_stmt(); + statement.setBytes(1, &job_id_bytes); + statement.setBytes(2, &client_id_bytes); + statement.addBatch(); + } + + // Tasks must be added in graph order to avoid the dangling reference. + std::vector const& input_task_ids = task_graph.get_input_tasks(); + absl::flat_hash_set heads; + for (boost::uuids::uuid const task_id : input_task_ids) { + heads.insert(task_id); + } + std::deque queue; + // First go over all heads + for (boost::uuids::uuid const task_id : heads) { + std::optional const task_option = task_graph.get_task(task_id); + if (!task_option.has_value()) { + static_cast(conn)->rollback(); + return StorageErr{ + StorageErrType::KeyNotFoundErr, + "Task graph inconsistent: head task not found" + }; + } + Task const* task = task_option.value(); + add_task_batch( + static_cast(batch), + job_id_bytes, + *task, + TaskState::Ready ); - sql::bytes task_id_bytes = uuid_get_bytes(task_id); - statement->setBytes(1, &task_id_bytes); - statement->executeUpdate(); + for (boost::uuids::uuid const id : task_graph.get_child_tasks(task_id)) { + std::vector const parents = task_graph.get_parent_tasks(id); + if (std::ranges::all_of(parents, [&](boost::uuids::uuid const& parent) { + return heads.contains(parent); + })) + { + queue.push_back(id); + } + } + } + // Then go over all tasks in queue + while (!queue.empty()) { + boost::uuids::uuid const task_id = queue.back(); + queue.pop_back(); + if (!heads.contains(task_id)) { + heads.insert(task_id); + std::optional const task_option = task_graph.get_task(task_id); + if (!task_option.has_value()) { + static_cast(conn)->rollback(); + return StorageErr{StorageErrType::KeyNotFoundErr, "Task graph inconsistent"}; + } + Task const* task = task_option.value(); + add_task_batch( + static_cast(batch), + job_id_bytes, + *task, + std::nullopt + ); + for (boost::uuids::uuid const id : task_graph.get_child_tasks(task_id)) { + std::vector const parents = task_graph.get_parent_tasks(id); + if (std::ranges::all_of(parents, [&](boost::uuids::uuid const& parent) { + return heads.contains(parent); + })) + { + queue.push_back(id); + } + } + } + } + + // Add all dependencies + for (std::pair const& pair : + task_graph.get_dependencies()) + { + sql::PreparedStatement& dep_statement + = static_cast(batch).get_task_dependency_stmt(); + sql::bytes parent_id_bytes = uuid_get_bytes(pair.first); + sql::bytes child_id_bytes = uuid_get_bytes(pair.second); + dep_statement.setBytes(1, &parent_id_bytes); + dep_statement.setBytes(2, &child_id_bytes); + dep_statement.addBatch(); + } + + // Add input tasks + for (size_t i = 0; i < input_task_ids.size(); i++) { + sql::PreparedStatement& input_statement + = static_cast(batch).get_input_task_stmt(); + input_statement.setBytes(1, &job_id_bytes); + sql::bytes task_id_bytes = uuid_get_bytes(input_task_ids[i]); + input_statement.setBytes(2, &task_id_bytes); + input_statement.setUInt(3, i); + input_statement.addBatch(); + } + // Add output tasks + std::vector const& output_task_ids = task_graph.get_output_tasks(); + for (size_t i = 0; i < output_task_ids.size(); i++) { + sql::PreparedStatement& output_statement + = static_cast(batch).get_output_task_stmt(); + output_statement.setBytes(1, &job_id_bytes); + sql::bytes task_id_bytes = uuid_get_bytes(output_task_ids[i]); + output_statement.setBytes(2, &task_id_bytes); + output_statement.setUInt(3, i); + output_statement.addBatch(); } } catch (sql::SQLException& e) { @@ -1126,7 +1146,7 @@ auto MySqlMetadataStorage::add_child( ) -> StorageErr { try { sql::bytes const job_id = uuid_get_bytes(child.get_id()); - add_task(static_cast(conn), job_id, child); + add_task(static_cast(conn), job_id, child, std::nullopt); // Add dependencies std::unique_ptr statement( @@ -1760,7 +1780,7 @@ auto MySqlMetadataStorage::set_scheduler_state( auto MySqlDataStorage::initialize(StorageConnection& conn) -> StorageErr { try { // Need to initialize metadata storage first so that foreign constraint is not voilated - for (char const* create_table_str : cCreateStorage) { + for (std::string const& create_table_str : mysql::cCreateStorage) { std::unique_ptr statement( static_cast(conn)->createStatement() ); diff --git a/src/spider/storage/MySqlStorage.hpp b/src/spider/storage/mysql/MySqlStorage.hpp similarity index 88% rename from src/spider/storage/MySqlStorage.hpp rename to src/spider/storage/mysql/MySqlStorage.hpp index a77953b14..b226e2b96 100644 --- a/src/spider/storage/MySqlStorage.hpp +++ b/src/spider/storage/mysql/MySqlStorage.hpp @@ -2,6 +2,7 @@ #define SPIDER_STORAGE_MYSQLSTORAGE_HPP #include +#include #include #include #include @@ -11,17 +12,19 @@ #include #include -#include "../core/Data.hpp" -#include "../core/Driver.hpp" -#include "../core/Error.hpp" -#include "../core/JobMetadata.hpp" -#include "../core/KeyValueData.hpp" -#include "../core/Task.hpp" -#include "../core/TaskGraph.hpp" -#include "DataStorage.hpp" -#include "MetadataStorage.hpp" +#include "../../core/Data.hpp" +#include "../../core/Driver.hpp" +#include "../../core/Error.hpp" +#include "../../core/JobMetadata.hpp" +#include "../../core/KeyValueData.hpp" +#include "../../core/Task.hpp" +#include "../../core/TaskGraph.hpp" +#include "../DataStorage.hpp" +#include "../JobSubmissionBatch.hpp" +#include "../MetadataStorage.hpp" +#include "../StorageConnection.hpp" #include "MySqlConnection.hpp" -#include "StorageConnection.hpp" +#include "MySqlJobSubmissionBatch.hpp" namespace spider::core { class MySqlMetadataStorage : public MetadataStorage { @@ -46,6 +49,13 @@ class MySqlMetadataStorage : public MetadataStorage { boost::uuids::uuid client_id, TaskGraph const& task_graph ) -> StorageErr override; + auto add_job_batch( + StorageConnection& conn, + JobSubmissionBatch& batch, + boost::uuids::uuid job_id, + boost::uuids::uuid client_id, + TaskGraph const& task_graph + ) -> StorageErr override; auto get_job_metadata(StorageConnection& conn, boost::uuids::uuid id, JobMetadata* job) -> StorageErr override; auto get_job_complete(StorageConnection& conn, boost::uuids::uuid id, bool* complete) @@ -123,7 +133,18 @@ class MySqlMetadataStorage : public MetadataStorage { private: std::string m_url; - static void add_task(MySqlConnection& conn, sql::bytes job_id, Task const& task); + static void add_task( + MySqlConnection& conn, + sql::bytes job_id, + Task const& task, + std::optional const& state + ); + static void add_task_batch( + MySqlJobSubmissionBatch& batch, + sql::bytes job_id, + Task const& task, + std::optional const& state + ); static auto fetch_full_task(MySqlConnection& conn, std::unique_ptr const& res) -> Task; }; diff --git a/src/spider/storage/mysql/mysql_stmt.hpp b/src/spider/storage/mysql/mysql_stmt.hpp new file mode 100644 index 000000000..db5c9f2b4 --- /dev/null +++ b/src/spider/storage/mysql/mysql_stmt.hpp @@ -0,0 +1,205 @@ +#ifndef SPIDER_STORAGE_MYSQLSTMT_HPP +#define SPIDER_STORAGE_MYSQLSTMT_HPP + +#include +#include + +namespace spider::core::mysql { +// NOLINTBEGIN(cert-err58-cpp) + +std::string const cCreateDriverTable = R"(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`) +))"; + +std::string const cCreateSchedulerTable = R"(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`) +))"; + +std::string const cCreateJobTable = R"(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`) +))"; + +std::string const cCreateTaskTable = R"(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`) +))"; + +std::string const cCreateInputTaskTable = R"(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`) +))"; + +std::string const cCreateOutputTaskTable = R"(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`) +))"; + +std::string const cCreateTaskInputTable = R"(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`) +))"; + +std::string const cCreateTaskOutputTable = R"(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`) +))"; + +std::string const cCreateTaskDependencyTable = R"(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 +))"; + +std::string const cCreateTaskInstanceTable = R"(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`) +))"; + +std::string const cCreateDataTable = R"(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`) +))"; + +std::string const cCreateDataLocalityTable = R"(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 +))"; + +std::string const cCreateDataRefDriverTable = R"(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 +))"; + +std::string const cCreateDataRefTaskTable = R"(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 +))"; + +std::string const cCreateClientKVDataTable = R"(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`) +))"; + +std::string const cCreateTaskKVDataTable = R"(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 +))"; + +std::array const cCreateStorage = { + cCreateDriverTable, // drivers table must be created before data_ref_driver + cCreateSchedulerTable, + cCreateJobTable, // jobs table must be created before task + cCreateTaskTable, // tasks table must be created before data_ref_task + cCreateDataTable, // data table must be created before task_outputs + cCreateDataLocalityTable, + cCreateDataRefDriverTable, + cCreateDataRefTaskTable, + cCreateClientKVDataTable, + cCreateTaskKVDataTable, + cCreateInputTaskTable, + cCreateOutputTaskTable, + cCreateTaskOutputTable, // task_outputs table must be created before task_inputs + cCreateTaskInputTable, + cCreateTaskDependencyTable, + cCreateTaskInstanceTable, +}; + +std::string const cInsertJob = R"(INSERT INTO `jobs` (`id`, `client_id`) VALUES (?, ?))"; + +std::string const cInsertTask + = R"(INSERT INTO `tasks` (`id`, `job_id`, `func_name`, `state`, `timeout`, `max_retry`) VALUES (?, ?, ?, ?, ?, ?))"; + +std::string const cInsertTaskInputOutput + = R"(INSERT INTO `task_inputs` (`task_id`, `position`, `type`, `output_task_id`, `output_task_position`) VALUES (?, ?, ?, ?, ?))"; + +std::string const cInsertTaskInputData + = R"(INSERT INTO `task_inputs` (`task_id`, `position`, `type`, `data_id`) VALUES (?, ?, ?, ?))"; + +std::string const cInsertTaskInputValue + = R"(INSERT INTO `task_inputs` (`task_id`, `position`, `type`, `value`) VALUES (?, ?, ?, ?))"; + +std::string const cInsertTaskOutput + = R"(INSERT INTO `task_outputs` (`task_id`, `position`, `type`) VALUES (?, ?, ?))"; + +std::string const cInsertTaskDependency + = R"(INSERT INTO `task_dependencies` (parent, child) VALUES (?, ?))"; + +std::string const cInsertInputTask + = R"(INSERT INTO `input_tasks` (`job_id`, `task_id`, `position`) VALUES (?, ?, ?))"; + +std::string const cInsertOutputTask + = R"(INSERT INTO `output_tasks` (`job_id`, `task_id`, `position`) VALUES (?, ?, ?))"; + +// NOLINTEND(cert-err58-cpp) +} // namespace spider::core::mysql + +#endif diff --git a/src/spider/worker/FunctionManager.hpp b/src/spider/worker/FunctionManager.hpp index 867764e30..9a226699e 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/MySqlConnection.hpp" +#include "../storage/mysql/MySqlConnection.hpp" #include "TaskExecutorMessage.hpp" // NOLINTBEGIN(cppcoreguidelines-macro-usage) diff --git a/src/spider/worker/WorkerClient.cpp b/src/spider/worker/WorkerClient.cpp index c22e52979..e59402245 100644 --- a/src/spider/worker/WorkerClient.cpp +++ b/src/spider/worker/WorkerClient.cpp @@ -23,7 +23,7 @@ #include "../scheduler/SchedulerMessage.hpp" #include "../storage/DataStorage.hpp" #include "../storage/MetadataStorage.hpp" -#include "../storage/MySqlConnection.hpp" +#include "../storage/mysql/MySqlConnection.hpp" namespace spider::worker { diff --git a/src/spider/worker/task_executor.cpp b/src/spider/worker/task_executor.cpp index 5799da2bd..3ff5deab9 100644 --- a/src/spider/worker/task_executor.cpp +++ b/src/spider/worker/task_executor.cpp @@ -24,7 +24,7 @@ #include "../io/MsgPack.hpp" // IWYU pragma: keep #include "../storage/DataStorage.hpp" #include "../storage/MetadataStorage.hpp" -#include "../storage/MySqlStorage.hpp" +#include "../storage/mysql/MySqlStorage.hpp" #include "DllLoader.hpp" #include "FunctionManager.hpp" #include "message_pipe.hpp" diff --git a/src/spider/worker/worker.cpp b/src/spider/worker/worker.cpp index cab609211..a99ae8ed4 100644 --- a/src/spider/worker/worker.cpp +++ b/src/spider/worker/worker.cpp @@ -38,8 +38,8 @@ #include "../io/Serializer.hpp" // IWYU pragma: keep #include "../storage/DataStorage.hpp" #include "../storage/MetadataStorage.hpp" -#include "../storage/MySqlConnection.hpp" -#include "../storage/MySqlStorage.hpp" +#include "../storage/mysql/MySqlConnection.hpp" +#include "../storage/mysql/MySqlStorage.hpp" #include "../utils/StopToken.hpp" #include "TaskExecutor.hpp" #include "WorkerClient.hpp" diff --git a/tests/client/client-test.cpp b/tests/client/client-test.cpp index 77fbcb3c5..8611bf19a 100644 --- a/tests/client/client-test.cpp +++ b/tests/client/client-test.cpp @@ -1,5 +1,6 @@ #include #include +#include #include #include @@ -39,6 +40,8 @@ auto parse_args(int const argc, char** argv) -> boost::program_options::variable constexpr int cCmdArgParseErr = 1; constexpr int cJobFailed = 2; +constexpr int cBatchSize = 10; + } // namespace // NOLINTNEXTLINE(bugprone-exception-escape) @@ -75,17 +78,17 @@ auto main(int argc, char** argv) -> int { spider::Data d1 = driver.get_data_builder().build(1); spider::Data d2 = driver.get_data_builder().build(2); spdlog::debug("Data created"); - spider::Job job = driver.start(graph, d1, d2, 3, 4); + spider::Job graph_job = driver.start(graph, d1, d2, 3, 4); spdlog::debug("Job started"); - job.wait_complete(); + graph_job.wait_complete(); spdlog::debug("Job completed"); - if (job.get_status() != spider::JobStatus::Succeeded) { + if (graph_job.get_status() != spider::JobStatus::Succeeded) { spdlog::error("Job failed"); return cJobFailed; } constexpr int cExpectedResult = 10; - if (job.get_result() != cExpectedResult) { - spdlog::error("Wrong job result. Get {}. Expect 10", job.get_result()); + if (graph_job.get_result() != cExpectedResult) { + spdlog::error("Wrong job result. Get {}. Expect 10", graph_job.get_result()); return cJobFailed; } @@ -157,5 +160,27 @@ auto main(int argc, char** argv) -> int { return cJobFailed; } + // Run batch submission + std::vector> jobs; + jobs.reserve(cBatchSize); + driver.begin_batch_start(); + for (int i = 0; i < cBatchSize; ++i) { + jobs.emplace_back(driver.start(&sum_test, i, i)); + } + driver.end_batch_start(); + for (int i = 0; i < cBatchSize; ++i) { + spider::Job& job = jobs[i]; + job.wait_complete(); + if (job.get_status() != spider::JobStatus::Succeeded) { + spdlog::error("Batch job failed"); + return cJobFailed; + } + int const result = job.get_result(); + if (result != i + i) { + spdlog::error("Batch job wrong result. Expect {}. Get {}.", i + i, result); + return cJobFailed; + } + } + return 0; } diff --git a/tests/scheduler/test-SchedulerPolicy.cpp b/tests/scheduler/test-SchedulerPolicy.cpp index 6f4b9005e..66d78a640 100644 --- a/tests/scheduler/test-SchedulerPolicy.cpp +++ b/tests/scheduler/test-SchedulerPolicy.cpp @@ -21,7 +21,7 @@ #include "../../src/spider/scheduler/FifoPolicy.hpp" #include "../../src/spider/storage/DataStorage.hpp" #include "../../src/spider/storage/MetadataStorage.hpp" -#include "../../src/spider/storage/MySqlConnection.hpp" +#include "../../src/spider/storage/mysql/MySqlConnection.hpp" #include "../storage/StorageTestHelper.hpp" namespace { diff --git a/tests/scheduler/test-SchedulerServer.cpp b/tests/scheduler/test-SchedulerServer.cpp index a1bc746c1..9233d3c1c 100644 --- a/tests/scheduler/test-SchedulerServer.cpp +++ b/tests/scheduler/test-SchedulerServer.cpp @@ -25,7 +25,7 @@ #include "../../src/spider/scheduler/SchedulerServer.hpp" #include "../../src/spider/storage/DataStorage.hpp" #include "../../src/spider/storage/MetadataStorage.hpp" -#include "../../src/spider/storage/MySqlConnection.hpp" +#include "../../src/spider/storage/mysql/MySqlConnection.hpp" #include "../../src/spider/utils/StopToken.hpp" #include "../storage/StorageTestHelper.hpp" diff --git a/tests/storage/StorageTestHelper.hpp b/tests/storage/StorageTestHelper.hpp index 09fc3e5a1..b79795e07 100644 --- a/tests/storage/StorageTestHelper.hpp +++ b/tests/storage/StorageTestHelper.hpp @@ -13,8 +13,8 @@ #include "../../src/spider/core/Error.hpp" #include "../../src/spider/storage/DataStorage.hpp" #include "../../src/spider/storage/MetadataStorage.hpp" -#include "../../src/spider/storage/MySqlConnection.hpp" -#include "../../src/spider/storage/MySqlStorage.hpp" +#include "../../src/spider/storage/mysql/MySqlConnection.hpp" +#include "../../src/spider/storage/mysql/MySqlStorage.hpp" namespace spider::test { char const* const cStorageUrl diff --git a/tests/storage/test-DataStorage.cpp b/tests/storage/test-DataStorage.cpp index 6bf250fd3..1e4a196c3 100644 --- a/tests/storage/test-DataStorage.cpp +++ b/tests/storage/test-DataStorage.cpp @@ -13,7 +13,7 @@ #include "../../src/spider/core/KeyValueData.hpp" #include "../../src/spider/core/Task.hpp" #include "../../src/spider/core/TaskGraph.hpp" -#include "../../src/spider/storage/MySqlConnection.hpp" +#include "../../src/spider/storage/mysql/MySqlConnection.hpp" #include "../utils/CoreDataUtils.hpp" #include "StorageTestHelper.hpp" diff --git a/tests/storage/test-MetadataStorage.cpp b/tests/storage/test-MetadataStorage.cpp index 39edf306c..9cc4612fc 100644 --- a/tests/storage/test-MetadataStorage.cpp +++ b/tests/storage/test-MetadataStorage.cpp @@ -18,7 +18,8 @@ #include "../../src/spider/core/Task.hpp" #include "../../src/spider/core/TaskGraph.hpp" #include "../../src/spider/storage/MetadataStorage.hpp" -#include "../../src/spider/storage/MySqlConnection.hpp" +#include "../../src/spider/storage/mysql/MySqlConnection.hpp" +#include "../../src/spider/storage/mysql/MySqlJobSubmissionBatch.hpp" #include "../utils/CoreTaskUtils.hpp" #include "StorageTestHelper.hpp" @@ -122,6 +123,136 @@ TEMPLATE_LIST_TEST_CASE( std::unique_ptr storage = spider::test::create_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); + spider::core::MySqlJobSubmissionBatch batch{conn}; + + boost::uuids::random_generator gen; + boost::uuids::uuid const job_id = gen(); + + // Create a complicated task graph + boost::uuids::uuid const client_id = gen(); + spider::core::Task child_task{"child"}; + spider::core::Task parent_1{"p1"}; + spider::core::Task parent_2{"p2"}; + parent_1.add_input(spider::core::TaskInput{"1", "float"}); + parent_1.add_input(spider::core::TaskInput{"2", "float"}); + parent_2.add_input(spider::core::TaskInput{"3", "int"}); + parent_2.add_input(spider::core::TaskInput{"4", "int"}); + parent_1.add_output(spider::core::TaskOutput{"float"}); + parent_2.add_output(spider::core::TaskOutput{"int"}); + child_task.add_input(spider::core::TaskInput{parent_1.get_id(), 0, "float"}); + child_task.add_input(spider::core::TaskInput{parent_2.get_id(), 0, "int"}); + child_task.add_output(spider::core::TaskOutput{"float"}); + spider::core::TaskGraph graph; + // Add task and dependencies to task graph in wrong order + graph.add_task(child_task); + graph.add_task(parent_1); + graph.add_task(parent_2); + graph.add_dependency(parent_2.get_id(), child_task.get_id()); + graph.add_dependency(parent_1.get_id(), child_task.get_id()); + graph.add_input_task(parent_1.get_id()); + graph.add_input_task(parent_2.get_id()); + graph.add_output_task(child_task.get_id()); + + // Get head tasks should succeed + std::vector heads = graph.get_input_tasks(); + REQUIRE(2 == heads.size()); + REQUIRE(heads[0] == parent_1.get_id()); + REQUIRE(heads[1] == parent_2.get_id()); + + std::chrono::system_clock::time_point const job_creation_time + = std::chrono::system_clock::now(); + + // Submit a simple job + boost::uuids::uuid const simple_job_id = gen(); + spider::core::Task const simple_task{"simple"}; + spider::core::TaskGraph simple_graph; + simple_graph.add_task(simple_task); + simple_graph.add_input_task(simple_task.get_id()); + simple_graph.add_output_task(simple_task.get_id()); + + heads = simple_graph.get_input_tasks(); + REQUIRE(1 == heads.size()); + 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); + + // 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(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(2 == job_ids.size()); + REQUIRE( + ((job_ids[0] == job_id && job_ids[1] == simple_job_id) + || (job_ids[0] == simple_job_id && job_ids[1] == job_id)) + ); + + // Get job metadata should get correct value + spider::core::JobMetadata job_metadata{}; + 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}; + // REQUIRE(job_creation_time + time_delta >= job_metadata.get_creation_time()); + // REQUIRE(job_creation_time - time_delta <= job_metadata.get_creation_time()); + + // Get task graph should succeed + spider::core::TaskGraph graph_res{}; + 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(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(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(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(2 == tasks.size()); + REQUIRE( + ((spider::test::task_equal(tasks[0], parent_1) + && spider::test::task_equal(tasks[1], parent_2)) + || (spider::test::task_equal(tasks[0], parent_2) + && spider::test::task_equal(tasks[1], parent_1))) + ); + + // Remove job should succeed + 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); + graph_res = spider::core::TaskGraph{}; + 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()); +} + +TEMPLATE_LIST_TEST_CASE( + "Job batch add, get and remove", + "[storage]", + spider::test::MetadataStorageTypeList +) { + std::unique_ptr storage + = spider::test::create_metadata_storage(); + std::variant conn_result = spider::core::MySqlConnection::create(storage->get_url()); REQUIRE(std::holds_alternative(conn_result)); diff --git a/tests/worker/test-FunctionManager.cpp b/tests/worker/test-FunctionManager.cpp index 2cd4de1fb..5f889bbde 100644 --- a/tests/worker/test-FunctionManager.cpp +++ b/tests/worker/test-FunctionManager.cpp @@ -17,7 +17,7 @@ #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/MySqlConnection.hpp" +#include "../../src/spider/storage/mysql/MySqlConnection.hpp" #include "../../src/spider/worker/FunctionManager.hpp" #include "../../src/spider/worker/FunctionNameManager.hpp" #include "../storage/StorageTestHelper.hpp" diff --git a/tests/worker/test-TaskExecutor.cpp b/tests/worker/test-TaskExecutor.cpp index 99cf0f445..dba48b7dd 100644 --- a/tests/worker/test-TaskExecutor.cpp +++ b/tests/worker/test-TaskExecutor.cpp @@ -22,7 +22,7 @@ #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/MySqlConnection.hpp" +#include "../../src/spider/storage/mysql/MySqlConnection.hpp" #include "../../src/spider/worker/FunctionManager.hpp" #include "../../src/spider/worker/TaskExecutor.hpp" #include "../storage/StorageTestHelper.hpp"