diff --git a/src/spider/storage/mysql/MySqlStorage.cpp b/src/spider/storage/mysql/MySqlStorage.cpp index 3170ff21b..776616628 100644 --- a/src/spider/storage/mysql/MySqlStorage.cpp +++ b/src/spider/storage/mysql/MySqlStorage.cpp @@ -1237,14 +1237,14 @@ auto MySqlMetadataStorage::get_ready_tasks( ) -> StorageErr { try { // Get all ready tasks from job that has not failed or cancelled - std::unique_ptr statement( + std::unique_ptr task_statement( static_cast(conn)->createStatement() ); - std::unique_ptr const res(statement->executeQuery( + std::unique_ptr const res{task_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(); @@ -1266,18 +1266,15 @@ auto MySqlMetadataStorage::get_ready_tasks( } // 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()); + std::unique_ptr job_statement{ + static_cast(conn)->createStatement() + }; + std::unique_ptr const job_res{job_statement->executeQuery( + "SELECT `id` , `client_id` , `creation_time` FROM `jobs` WHERE `id` IN (SELECT " + "DISTINCT `job_id` FROM `tasks` WHERE `state` = 'ready' AND `job_id` NOT IN " + "(SELECT `job_id` FROM `tasks` WHERE `state` = 'fail' OR `state` = 'cancel'))" + )}; + 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")); @@ -1300,22 +1297,18 @@ auto MySqlMetadataStorage::get_ready_tasks( } // 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()); + std::unique_ptr locality_statement{ + static_cast(conn)->createStatement() + }; + std::unique_ptr const locality_res{locality_statement->executeQuery( + "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` IN (SELECT `id` " + "FROM `tasks` WHERE `state` = 'ready' AND `job_id` NOT IN (SELECT `job_id` " + "FROM `tasks` WHERE `state` = 'fail' OR `state` = 'cancel'))" + )}; + 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"); @@ -1328,8 +1321,14 @@ auto MySqlMetadataStorage::get_ready_tasks( } // Add all tasks to the output - for (auto const& ite : new_tasks) { - tasks->emplace_back(ite.second); + absl::flat_hash_set task_ids; + for (ScheduleTaskMetadata const& task : *tasks) { + task_ids.insert(task.get_id()); + } + for (auto const& [task_id, task] : new_tasks) { + if (task_ids.find(task_id) == task_ids.end()) { + tasks->emplace_back(task); + } } } catch (sql::SQLException& e) { static_cast(conn)->rollback(); @@ -1629,64 +1628,21 @@ auto MySqlMetadataStorage::get_task_timeout( std::vector* tasks ) -> StorageErr { try { - std::unique_ptr statement( + std::unique_ptr task_statement( static_cast(conn)->createStatement() ); - 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 `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`, `job_id` FROM `tasks` WHERE `id` = ?" - ) - ); - for (boost::uuids::uuid const& task_id : task_ids) { - sql::bytes task_id_bytes = uuid_get_bytes(task_id); - task_statement->setBytes(1, &task_id_bytes); - task_statement->addBatch(); - } - task_statement->execute(); - std::unique_ptr const task_res(task_statement->getResultSet()); + std::unique_ptr const task_res{task_statement->executeQuery( + "SELECT `id`, `func_name`, `job_id` FROM `tasks` WHERE `id` IN (SELECT " + "`tasks`.`id` FROM `task_instances` JOIN `tasks` ON `task_instances`.`task_id` = " + "`tasks`.`id` WHERE `tasks`.`timeout` > 0.0001 AND `tasks`.`state` = 'running' AND " + "TIMESTAMPDIFF(MICROSECOND, `task_instances`.`start_time`, CURRENT_TIMESTAMP()) > " + "`tasks`.`timeout` * 1000) AND `id` NOT IN (SELECT `tasks`.`id` FROM " + "`task_instances` JOIN `tasks` ON `task_instances`.`task_id` = `tasks`.`id` WHERE " + "`tasks`.`timeout` > 0.0001 AND `tasks`.`state` = 'running' AND " + "TIMESTAMPDIFF(MICROSECOND, `task_instances`.`start_time`, CURRENT_TIMESTAMP()) < " + "`tasks`.` timeout` * 1000)" + )}; absl::flat_hash_map new_tasks; absl::flat_hash_map> job_id_to_task_ids; @@ -1703,18 +1659,22 @@ auto MySqlMetadataStorage::get_task_timeout( } // 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()); + std::unique_ptr job_statement{ + static_cast(conn)->createStatement() + }; + std::unique_ptr const job_res{job_statement->executeQuery( + "SELECT `jobs`.`id`, `jobs`.`client_id`, `jobs`.`creation_time` FROM `jobs` JOIN " + "`tasks` ON `jobs`.`id` = `tasks`.`job_id` WHERE `tasks`.`id` IN (SELECT " + "`tasks`.`id` FROM `task_instances` JOIN `tasks` ON `task_instances`.`task_id` = " + "`tasks`.`id` WHERE `tasks`.`timeout` > 0.0001 AND `tasks`.`state` = 'running' AND " + "TIMESTAMPDIFF(MICROSECOND, `task_instances`.`start_time`, CURRENT_TIMESTAMP()) > " + "`tasks`.`timeout` * 1000) AND `tasks`.`id` NOT IN (SELECT `tasks`.`id` FROM " + "`task_instances` JOIN `tasks` ON `task_instances`.`task_id` = `tasks`.`id` WHERE " + "`tasks`.`timeout` > 0.0001 AND `tasks`.`state` = 'running' AND " + "TIMESTAMPDIFF(MICROSECOND, `task_instances`.`start_time`, CURRENT_TIMESTAMP()) < " + "`tasks`.` timeout` * 1000)" + )}; + 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")); @@ -1737,22 +1697,23 @@ auto MySqlMetadataStorage::get_task_timeout( } // 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()); + std::unique_ptr locality_statement{ + static_cast(conn)->createStatement() + }; + std::unique_ptr const locality_res{locality_statement->executeQuery( + "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` IN (SELECT `tasks`.`id` FROM `task_instances` JOIN " + "`tasks` ON `task_instances`.`task_id` = `tasks`.`id` WHERE `tasks`.`timeout` > " + "0.0001 AND `tasks`.`state` = 'running' AND TIMESTAMPDIFF(MICROSECOND, " + "`task_instances`.`start_time`, CURRENT_TIMESTAMP()) > `tasks`.`timeout` * 1000) " + "AND `task_inputs`.`task_id` NOT IN (SELECT `tasks`.`id` FROM `task_instances` " + "JOIN `tasks` ON `task_instances`.`task_id` = `tasks`.`id` WHERE `tasks`.`timeout` " + "> 0.0001 AND `tasks`.`state` = 'running' AND TIMESTAMPDIFF(MICROSECOND, " + "`task_instances`.`start_time`, CURRENT_TIMESTAMP()) < `tasks`.` timeout` * 1000)" + )}; + 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"); @@ -1765,8 +1726,14 @@ auto MySqlMetadataStorage::get_task_timeout( } // Add all tasks to the output - for (auto const& iter : new_tasks) { - tasks->emplace_back(iter.second); + absl::flat_hash_set task_ids; + for (ScheduleTaskMetadata const& task : *tasks) { + task_ids.insert(task.get_id()); + } + for (auto const& [task_id, task] : new_tasks) { + if (task_ids.find(task_id) == task_ids.end()) { + tasks->emplace_back(task); + } } } catch (sql::SQLException& e) { static_cast(conn)->rollback();