diff --git a/src/spider/core/Task.hpp b/src/spider/core/Task.hpp index 22016738e..368432a89 100644 --- a/src/spider/core/Task.hpp +++ b/src/spider/core/Task.hpp @@ -1,6 +1,7 @@ #ifndef SPIDER_CORE_TASK_HPP #define SPIDER_CORE_TASK_HPP +#include #include #include #include @@ -109,6 +110,64 @@ enum class TaskState : std::uint8_t { Canceled, }; +class ScheduleTaskMetadata { +public: + ScheduleTaskMetadata( + boost::uuids::uuid id, + std::string function_name, + boost::uuids::uuid job_id + ) + : m_id(id), + m_function_name(std::move(function_name)), + m_job_id(job_id) {} + + ScheduleTaskMetadata() = default; + + [[nodiscard]] auto get_id() const -> boost::uuids::uuid { return m_id; } + + [[nodiscard]] auto get_function_name() const -> std::string const& { return m_function_name; } + + [[nodiscard]] auto get_job_id() const -> boost::uuids::uuid { return m_job_id; } + + [[nodiscard]] auto get_client_id() const -> boost::uuids::uuid { return m_client_id; } + + [[nodiscard]] auto get_job_creation_time() const -> std::chrono::system_clock::time_point { + return m_job_creation_time; + } + + [[nodiscard]] auto get_hard_localities() const -> std::vector const& { + return m_hard_localities; + } + + [[nodiscard]] auto get_soft_localities() const -> std::vector const& { + return m_soft_localities; + } + + auto set_client_id(boost::uuids::uuid const client_id) -> void { m_client_id = client_id; } + + auto set_job_creation_time(std::chrono::system_clock::time_point const job_creation_time + ) -> void { + m_job_creation_time = job_creation_time; + } + + auto add_hard_locality(std::string const& locality) -> void { + m_hard_localities.push_back(locality); + } + + auto add_soft_locality(std::string const& locality) -> void { + m_soft_localities.push_back(locality); + } + +private: + boost::uuids::uuid m_id; + std::string m_function_name; + boost::uuids::uuid m_job_id; + boost::uuids::uuid m_client_id; + std::chrono::system_clock::time_point m_job_creation_time; + std::vector m_hard_localities; + std::vector m_soft_localities; +}; + class Task { public: explicit Task(std::string function_name) : m_function_name(std::move(function_name)) { diff --git a/src/spider/scheduler/FifoPolicy.cpp b/src/spider/scheduler/FifoPolicy.cpp index cb93fab48..ae57dd60a 100644 --- a/src/spider/scheduler/FifoPolicy.cpp +++ b/src/spider/scheduler/FifoPolicy.cpp @@ -1,22 +1,14 @@ #include "FifoPolicy.hpp" #include -#include #include #include #include -#include #include -#include #include -#include #include -#include -#include -#include "../core/Data.hpp" -#include "../core/JobMetadata.hpp" #include "../core/Task.hpp" #include "../storage/DataStorage.hpp" #include "../storage/MetadataStorage.hpp" @@ -24,42 +16,6 @@ namespace spider::scheduler { -auto FifoPolicy::task_locality_satisfied(spider::core::Task const& task, std::string const& addr) - -> bool { - for (auto const& input : task.get_inputs()) { - if (input.get_value().has_value()) { - continue; - } - std::optional optional_data_id = input.get_data_id(); - if (!optional_data_id.has_value()) { - continue; - } - boost::uuids::uuid const data_id = optional_data_id.value(); - core::Data data; - 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()) { - throw std::runtime_error( - fmt::format("Data with id {} not exists.", to_string((data_id))) - ); - } - m_data_cache.emplace(data_id, data); - } - if (false == data.is_hard_locality()) { - continue; - } - std::vector const& locality = data.get_locality(); - if (locality.empty()) { - continue; - } - if (std::ranges::find(locality, addr) == locality.end()) { - return false; - } - } - return true; -} - FifoPolicy::FifoPolicy( std::shared_ptr const& metadata_store, std::shared_ptr const& data_store, @@ -81,56 +37,35 @@ auto FifoPolicy::schedule_next( } auto const reverse_begin = std::reverse_iterator(m_tasks.end()); auto const reverse_end = std::reverse_iterator(m_tasks.begin()); - auto const it = std::find_if(reverse_begin, reverse_end, [&](core::Task const& task) { - return task_locality_satisfied(task, worker_addr); - }); + auto const it + = std::find_if(reverse_begin, reverse_end, [&](core::ScheduleTaskMetadata const& task) { + std::vector const& hard_localities = task.get_hard_localities(); + if (hard_localities.empty()) { + return true; + } + // If the worker address is in the hard localities, then the task can be + // scheduled. + return std::ranges::find(hard_localities, worker_addr) != hard_localities.end(); + }); if (it == reverse_end) { return std::nullopt; } boost::uuids::uuid const task_id = it->get_id(); - for (core::TaskInput const& input : it->get_inputs()) { - std::optional const data_id = input.get_data_id(); - if (data_id.has_value()) { - m_data_cache.erase(data_id.value()); - } - } m_tasks.erase(std::next(it).base()); return task_id; } auto FifoPolicy::fetch_tasks() -> void { - m_data_cache.clear(); m_metadata_store->get_ready_tasks(*m_conn, &m_tasks); - std::vector> instances; - m_metadata_store->get_task_timeout(*m_conn, &instances); - for (auto const& [instance, task] : instances) { - m_tasks.emplace_back(task); - } + m_metadata_store->get_task_timeout(*m_conn, &m_tasks); // Sort tasks based on job creation time in descending order. - // NOLINTNEXTLINE(misc-include-cleaner) - absl::flat_hash_map> - job_metadata_map; - 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()) { - throw std::runtime_error(fmt::format("Task with id {} not exists.", to_string(task_id)) - ); - } - if (job_metadata_map.contains(job_id)) { - 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()) { - throw std::runtime_error(fmt::format("Job with id {} not exists.", to_string(job_id))); - } - job_metadata_map[job_id] = job_metadata; - return job_metadata.get_creation_time(); - }; - std::ranges::sort(m_tasks, [&](core::Task const& a, core::Task const& b) { - return get_task_job_creation_time(a.get_id()) > get_task_job_creation_time(b.get_id()); - }); + std::ranges::sort( + m_tasks, + [&](core::ScheduleTaskMetadata const& a, core::ScheduleTaskMetadata const& b) { + return a.get_job_creation_time() > b.get_job_creation_time(); + } + ); } } // namespace spider::scheduler diff --git a/src/spider/scheduler/FifoPolicy.hpp b/src/spider/scheduler/FifoPolicy.hpp index 0629128de..2c142f1f6 100644 --- a/src/spider/scheduler/FifoPolicy.hpp +++ b/src/spider/scheduler/FifoPolicy.hpp @@ -6,7 +6,6 @@ #include #include -#include #include #include "../core/Task.hpp" @@ -30,15 +29,12 @@ class FifoPolicy final : public SchedulerPolicy { private: auto fetch_tasks() -> void; - auto task_locality_satisfied(core::Task const& task, std::string const& addr) -> bool; std::shared_ptr m_metadata_store; std::shared_ptr m_data_store; std::shared_ptr m_conn; - std::vector m_tasks; - // NOLINTNEXTLINE(misc-include-cleaner) - absl::flat_hash_map> m_data_cache; + std::vector m_tasks; }; } // namespace spider::scheduler diff --git a/src/spider/storage/MetadataStorage.hpp b/src/spider/storage/MetadataStorage.hpp index 7e7645eb8..24f3ec7b8 100644 --- a/src/spider/storage/MetadataStorage.hpp +++ b/src/spider/storage/MetadataStorage.hpp @@ -2,7 +2,6 @@ #define SPIDER_STORAGE_METADATASTORAGE_HPP #include -#include #include #include @@ -78,8 +77,8 @@ class MetadataStorage { boost::uuids::uuid id, boost::uuids::uuid* job_id ) -> StorageErr = 0; - virtual auto get_ready_tasks(StorageConnection& conn, std::vector* tasks) -> StorageErr - = 0; + virtual auto get_ready_tasks(StorageConnection& conn, std::vector* tasks) + -> StorageErr = 0; virtual auto set_task_state(StorageConnection& conn, boost::uuids::uuid id, TaskState state) -> StorageErr = 0; virtual auto set_task_running(StorageConnection& conn, boost::uuids::uuid id) -> StorageErr = 0; @@ -98,10 +97,8 @@ class MetadataStorage { TaskInstance const& instance, std::string const& error ) -> StorageErr = 0; - virtual auto get_task_timeout( - StorageConnection& conn, - std::vector>* tasks - ) -> StorageErr = 0; + virtual auto get_task_timeout(StorageConnection& conn, std::vector* tasks) + -> StorageErr = 0; virtual auto get_child_tasks( StorageConnection& conn, boost::uuids::uuid id, diff --git a/src/spider/storage/mysql/MySqlStorage.cpp b/src/spider/storage/mysql/MySqlStorage.cpp index c71c242a4..3170ff21b 100644 --- a/src/spider/storage/mysql/MySqlStorage.cpp +++ b/src/spider/storage/mysql/MySqlStorage.cpp @@ -14,6 +14,7 @@ #include #include +#include #include #include #include @@ -1230,20 +1231,105 @@ auto MySqlMetadataStorage::get_task_job_id( return StorageErr{}; } -auto MySqlMetadataStorage::get_ready_tasks(StorageConnection& conn, std::vector* tasks) - -> StorageErr { +auto MySqlMetadataStorage::get_ready_tasks( + StorageConnection& conn, + std::vector* tasks +) -> StorageErr { try { // Get all ready tasks from job that has not failed or cancelled std::unique_ptr statement( static_cast(conn)->createStatement() ); - std::unique_ptr res(statement->executeQuery( - "SELECT `id`, `func_name`, `state`, `timeout` FROM `tasks` WHERE `state` = 'ready' " + std::unique_ptr const res(statement->executeQuery( + "SELECT `id`, `func_name`, `job_id` FROM `tasks` WHERE `state` = 'ready' " "AND `job_id` NOT IN (SELECT `job_id` FROM `tasks` WHERE `state` = 'fail' OR " "`state` = 'cancel')" )); + + if (res->rowsCount() == 0) { + static_cast(conn)->commit(); + return StorageErr{}; + } + + absl::flat_hash_map new_tasks; + absl::flat_hash_map> job_id_to_task_ids; while (res->next()) { - tasks->emplace_back(fetch_full_task(static_cast(conn), res)); + boost::uuids::uuid const task_id = read_id(res->getBinaryStream("id")); + boost::uuids::uuid const job_id = read_id(res->getBinaryStream("job_id")); + std::string const function_name = get_sql_string(res->getString("func_name")); + new_tasks.emplace(task_id, ScheduleTaskMetadata{task_id, function_name, job_id}); + if (job_id_to_task_ids.find(job_id) == job_id_to_task_ids.end()) { + job_id_to_task_ids[job_id] = std::vector{task_id}; + } else { + job_id_to_task_ids[job_id].emplace_back(task_id); + } + } + + // Get all job metadata + std::unique_ptr job_statement( + static_cast(conn)->prepareStatement( + "SELECT `id`, `client_id`, `creation_time` FROM `jobs` WHERE `id` = ?" + ) + ); + for (auto const& iter : job_id_to_task_ids) { + sql::bytes job_id_bytes = uuid_get_bytes(iter.first); + job_statement->setBytes(1, &job_id_bytes); + job_statement->addBatch(); + } + job_statement->execute(); + std::unique_ptr const job_res(job_statement->getResultSet()); + while (job_res->next()) { + boost::uuids::uuid const job_id = read_id(job_res->getBinaryStream("id")); + boost::uuids::uuid const client_id = read_id(job_res->getBinaryStream("client_id")); + std::optional const optional_creation_time + = parse_timestamp(get_sql_string(job_res->getString("creation_time"))); + if (false == optional_creation_time.has_value()) { + static_cast(conn)->rollback(); + return StorageErr{ + StorageErrType::OtherErr, + fmt::format( + "Cannot parse timestamp {}", + get_sql_string(job_res->getString("creation_time")) + ) + }; + } + for (boost::uuids::uuid const& task_id : job_id_to_task_ids[job_id]) { + new_tasks[task_id].set_client_id(client_id); + new_tasks[task_id].set_job_creation_time(optional_creation_time.value()); + } + } + + // Get all data localities + std::unique_ptr locality_statement( + static_cast(conn)->prepareStatement( + "SELECT `task_inputs`.`task_id`, `data`.`hard_locality`, " + "`data_locality`.`address` FROM `task_inputs` JOIN `data` ON " + "`task_inputs`.`data_id` = `data`.`id` JOIN `data_locality` ON `data`.`id` " + "= `data_locality`.`id` WHERE `task_inputs`.`task_id` = ? AND " + "`task_inputs`.`task_id` IS NOT NULL" + ) + ); + for (auto const& iter : new_tasks) { + sql::bytes task_id_bytes = uuid_get_bytes(iter.first); + locality_statement->setBytes(1, &task_id_bytes); + locality_statement->addBatch(); + } + locality_statement->execute(); + std::unique_ptr const locality_res(locality_statement->getResultSet()); + while (locality_res->next()) { + boost::uuids::uuid const task_id = read_id(locality_res->getBinaryStream("task_id")); + bool const hard_locality = locality_res->getBoolean("hard_locality"); + std::string const address = get_sql_string(locality_res->getString("address")); + if (hard_locality) { + new_tasks[task_id].add_hard_locality(address); + } else { + new_tasks[task_id].add_soft_locality(address); + } + } + + // Add all tasks to the output + for (auto const& ite : new_tasks) { + tasks->emplace_back(ite.second); } } catch (sql::SQLException& e) { static_cast(conn)->rollback(); @@ -1540,50 +1626,148 @@ auto MySqlMetadataStorage::task_fail( auto MySqlMetadataStorage::get_task_timeout( StorageConnection& conn, - std::vector>* tasks + std::vector* tasks ) -> StorageErr { try { std::unique_ptr statement( static_cast(conn)->createStatement() ); - std::unique_ptr res(statement->executeQuery( - "SELECT `t1`.`id`, `t1`.`task_id` FROM `task_instances` as `t1` JOIN `tasks` ON " + std::unique_ptr const task_timeout_res(statement->executeQuery( + "SELECT `t1`.`task_id` FROM `task_instances` as `t1` JOIN `tasks` ON " "`t1`.`task_id` = `tasks`.`id` WHERE `tasks`.`timeout` > 0.0001 AND " "TIMESTAMPDIFF(MICROSECOND, `t1`.`start_time`, CURRENT_TIMESTAMP()) > " "`tasks`.`timeout` * 1000" )); + if (task_timeout_res->rowsCount() == 0) { + static_cast(conn)->commit(); + return StorageErr{}; + } + std::unique_ptr not_timeout_statement( static_cast(conn)->prepareStatement( - "SELECT FROM `task_instances` as `t1` JOIN `tasks` ON`t1`.`task_id` = " - "`tasks`.`id` WHERE `t1.task_id` = ? AND TIMESTAMPDIFF(MICROSECOND, " - "`t1`.`start_time`, CURRENT_TIMESTAMP()) < `tasks`.`timeout` * 1000" + "SELECT `t1`.`task_id` FROM `task_instances` as `t1` JOIN `tasks` ON " + "`t1`.`task_id` = `tasks`.`id` WHERE `t1`.`task_id` = ? AND " + "TIMESTAMPDIFF(MICROSECOND, `t1`.`start_time`, CURRENT_TIMESTAMP()) < " + "`tasks`.`timeout` * 1000" ) ); + absl::flat_hash_set task_ids; + while (task_timeout_res->next()) { + boost::uuids::uuid const task_id + = read_id(task_timeout_res->getBinaryStream("task_id")); + task_ids.insert(task_id); + sql::bytes task_id_bytes = uuid_get_bytes(task_id); + not_timeout_statement->setBytes(1, &task_id_bytes); + not_timeout_statement->addBatch(); + } + not_timeout_statement->execute(); + std::unique_ptr const not_timeout_res(not_timeout_statement->getResultSet() + ); + while (not_timeout_res->next()) { + boost::uuids::uuid const task_id = read_id(not_timeout_res->getBinaryStream("task_id")); + task_ids.erase(task_id); + } + + if (task_ids.empty()) { + static_cast(conn)->commit(); + return StorageErr{}; + } + + // Get task metadata std::unique_ptr task_statement( static_cast(conn)->prepareStatement( - "SELECT `id`, `func_name`, `state`, `timeout` FROM `tasks` WHERE `id` = ?" + "SELECT `id`, `func_name`, `job_id` FROM `tasks` WHERE `id` = ?" ) ); - while (res->next()) { - boost::uuids::uuid const task_instance_id = read_id(res->getBinaryStream("id")); - boost::uuids::uuid const task_id = read_id(res->getBinaryStream("task_id")); + for (boost::uuids::uuid const& task_id : task_ids) { sql::bytes task_id_bytes = uuid_get_bytes(task_id); - // Check all task instance have timed out - not_timeout_statement->setBytes(1, &task_id_bytes); - std::unique_ptr not_timeout_res(not_timeout_statement->executeQuery()); - if (not_timeout_res->rowsCount() > 0) { - continue; + task_statement->setBytes(1, &task_id_bytes); + task_statement->addBatch(); + } + task_statement->execute(); + std::unique_ptr const task_res(task_statement->getResultSet()); + + absl::flat_hash_map new_tasks; + absl::flat_hash_map> job_id_to_task_ids; + while (task_res->next()) { + boost::uuids::uuid const task_id = read_id(task_res->getBinaryStream("id")); + boost::uuids::uuid const job_id = read_id(task_res->getBinaryStream("job_id")); + std::string const function_name = get_sql_string(task_res->getString("func_name")); + new_tasks.emplace(task_id, ScheduleTaskMetadata{task_id, function_name, job_id}); + if (job_id_to_task_ids.find(job_id) == job_id_to_task_ids.end()) { + job_id_to_task_ids[job_id] = std::vector{task_id}; + } else { + job_id_to_task_ids[job_id].emplace_back(task_id); } + } - // Fetch task - task_statement->setBytes(1, &task_id_bytes); - std::unique_ptr task_res(task_statement->executeQuery()); - if (task_res->next()) { - Task const task = fetch_full_task(static_cast(conn), task_res); - tasks->emplace_back(TaskInstance{task_instance_id, task_id}, task); + // Get all job metadata + std::unique_ptr job_statement( + static_cast(conn)->prepareStatement( + "SELECT `id`, `client_id`, `creation_time` FROM `jobs` WHERE `id` = ?" + ) + ); + for (auto const& iter : job_id_to_task_ids) { + sql::bytes job_id_bytes = uuid_get_bytes(iter.first); + job_statement->setBytes(1, &job_id_bytes); + job_statement->addBatch(); + } + job_statement->execute(); + std::unique_ptr const job_res(job_statement->getResultSet()); + while (job_res->next()) { + boost::uuids::uuid const job_id = read_id(job_res->getBinaryStream("id")); + boost::uuids::uuid const client_id = read_id(job_res->getBinaryStream("client_id")); + std::optional const optional_creation_time + = parse_timestamp(get_sql_string(job_res->getString("creation_time"))); + if (false == optional_creation_time.has_value()) { + static_cast(conn)->rollback(); + return StorageErr{ + StorageErrType::OtherErr, + fmt::format( + "Cannot parse timestamp {}", + get_sql_string(job_res->getString("creation_time")) + ) + }; + } + for (boost::uuids::uuid const& task_id : job_id_to_task_ids[job_id]) { + new_tasks[task_id].set_client_id(client_id); + new_tasks[task_id].set_job_creation_time(optional_creation_time.value()); + } + } + + // Get all data localities + std::unique_ptr locality_statement( + static_cast(conn)->prepareStatement( + "SELECT `task_inputs`.`task_id`, `data`.`hard_locality`, " + "`data_locality`.`address` FROM `task_inputs` JOIN `data` ON " + "`task_inputs`.`data_id` = `data`.`id` JOIN `data_locality` ON `data`.`id` " + "= `data_locality`.`id` WHERE `task_inputs`.`task_id` = ? AND " + "`task_inputs`.`task_id` IS NOT NULL" + ) + ); + for (auto const& iter : new_tasks) { + sql::bytes task_id_bytes = uuid_get_bytes(iter.first); + locality_statement->setBytes(1, &task_id_bytes); + locality_statement->addBatch(); + } + locality_statement->execute(); + std::unique_ptr const locality_res(locality_statement->getResultSet()); + while (locality_res->next()) { + boost::uuids::uuid const task_id = read_id(locality_res->getBinaryStream("task_id")); + bool const hard_locality = locality_res->getBoolean("hard_locality"); + std::string const address = get_sql_string(locality_res->getString("address")); + if (hard_locality) { + new_tasks[task_id].add_hard_locality(address); + } else { + new_tasks[task_id].add_soft_locality(address); } } + + // Add all tasks to the output + for (auto const& iter : new_tasks) { + tasks->emplace_back(iter.second); + } } catch (sql::SQLException& e) { static_cast(conn)->rollback(); return StorageErr{StorageErrType::OtherErr, e.what()}; diff --git a/src/spider/storage/mysql/MySqlStorage.hpp b/src/spider/storage/mysql/MySqlStorage.hpp index fc12c3658..a563b54c5 100644 --- a/src/spider/storage/mysql/MySqlStorage.hpp +++ b/src/spider/storage/mysql/MySqlStorage.hpp @@ -4,7 +4,6 @@ #include #include #include -#include #include #include @@ -81,7 +80,8 @@ class MySqlMetadataStorage : public MetadataStorage { get_task(StorageConnection& conn, boost::uuids::uuid id, Task* task) -> StorageErr override; auto get_task_job_id(StorageConnection& conn, boost::uuids::uuid id, boost::uuids::uuid* job_id) -> StorageErr override; - auto get_ready_tasks(StorageConnection& conn, std::vector* tasks) -> StorageErr override; + auto get_ready_tasks(StorageConnection& conn, std::vector* tasks) + -> StorageErr override; auto set_task_state(StorageConnection& conn, boost::uuids::uuid id, TaskState state) -> StorageErr override; auto set_task_running(StorageConnection& conn, boost::uuids::uuid id) -> StorageErr override; @@ -96,10 +96,8 @@ class MySqlMetadataStorage : public MetadataStorage { ) -> StorageErr override; auto task_fail(StorageConnection& conn, TaskInstance const& instance, std::string const& error) -> StorageErr override; - auto get_task_timeout( - StorageConnection& conn, - std::vector>* tasks - ) -> StorageErr override; + auto get_task_timeout(StorageConnection& conn, std::vector* tasks) + -> StorageErr override; auto get_child_tasks( StorageConnection& conn, boost::uuids::uuid id,