diff --git a/mindspore/ccsrc/armour/CMakeLists.txt b/mindspore/ccsrc/armour/CMakeLists.txt new file mode 100644 index 00000000000..c4e0cfdd05d --- /dev/null +++ b/mindspore/ccsrc/armour/CMakeLists.txt @@ -0,0 +1,12 @@ +file(GLOB_RECURSE ARMOUR_FILES RELATIVE ${CMAKE_CURRENT_SOURCE_DIR} "*.cc") + +set(SERVER_FLATBUFFER_OUTPUT "${CMAKE_BINARY_DIR}/schema") +set(FBS_FILES + ${CMAKE_CURRENT_SOURCE_DIR}/../../schema/cipher.fbs + ${CMAKE_CURRENT_SOURCE_DIR}/../../schema/fl_job.fbs + ) + +set_property(SOURCE ${ARMOUR_FILES} PROPERTY COMPILE_DEFINITIONS SUBMODULE_ID=mindspore::SubModuleId::SM_ARMOUR) +add_library(_mindspore_armour_obj OBJECT ${ARMOUR_FILES}) +add_dependencies(_mindspore_armour_obj generated_fbs_files) +target_link_libraries(_mindspore_armour_obj mindspore::flatbuffers) \ No newline at end of file diff --git a/mindspore/ccsrc/armour/cipher/cipher_init.cc b/mindspore/ccsrc/armour/cipher/cipher_init.cc new file mode 100644 index 00000000000..81f27b414c7 --- /dev/null +++ b/mindspore/ccsrc/armour/cipher/cipher_init.cc @@ -0,0 +1,82 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "armour/cipher/cipher_init.h" +#include "ps/server/common.h" +#include "armour/cipher/cipher_meta_storage.h" + +namespace mindspore { +namespace armour { + +bool CipherInit::Init(const CipherPublicPara ¶m, size_t time_out_mutex, size_t cipher_initial_client_cnt, + size_t cipher_exchange_secrets_cnt, size_t cipher_share_secrets_cnt, + size_t cipher_get_clientlist_cnt, size_t cipher_reconstruct_secrets_down_cnt, + size_t cipher_reconstruct_secrets_up_cnt) { + MS_LOG(INFO) << "CipherInit::Init START"; + int return_num = 0; + cipher_meta_storage_.RegisterClass(); + const std::string new_prime(reinterpret_cast(param.prime), PRIME_MAX_LEN); + cipher_meta_storage_.RegisterPrime(ps::server::kCtxCipherPrimer, new_prime); + if (!cipher_meta_storage_.GetPrimeFromServer(ps::server::kCtxCipherPrimer, publicparam_.prime)) { + MS_LOG(ERROR) << "Cipher Param Update is invalid."; + return false; + } + + return_num = memcpy_s(publicparam_.p, SECRET_MAX_LEN, param.p, SECRET_MAX_LEN); + if (return_num != 0) { + return false; + } + publicparam_.g = param.g; + publicparam_.t = param.t; + secrets_minnums_ = param.t; + client_num_need_ = cipher_initial_client_cnt; + featuremap_ = 1000; // todo: wait for other code + // merge.ps::server::DistributedMetadataStore::GetInstance().model_size() / sizeof(float); + share_clients_num_need_ = cipher_share_secrets_cnt; + reconstruct_clients_num_need_ = cipher_reconstruct_secrets_down_cnt + 1; + get_model_num_need_ = cipher_get_clientlist_cnt; + time_out_mutex_ = time_out_mutex; + publicparam_.dp_eps = param.dp_eps; + publicparam_.dp_delta = param.dp_delta; + publicparam_.dp_norm_clip = param.dp_norm_clip; + publicparam_.encrypt_type = param.encrypt_type; + + MS_LOG(INFO) << " CipherInit client_num_need_ : " << client_num_need_; + MS_LOG(INFO) << " CipherInit share_clients_num_need_ : " << share_clients_num_need_; + MS_LOG(INFO) << " CipherInit reconstruct_clients_num_need_ : " << reconstruct_clients_num_need_; + MS_LOG(INFO) << " CipherInit get_model_num_need_ : " << get_model_num_need_; + MS_LOG(INFO) << " CipherInit featuremap_ : " << featuremap_; + + if (Check_Parames() == false) { + MS_LOG(ERROR) << "Cipher parameters are illegal."; + return false; + } + + MS_LOG(INFO) << " CipherInit::Init Success"; + return true; +} + +bool CipherInit::Check_Parames() { + if (featuremap_ < 1 || secrets_minnums_ < 1 || share_clients_num_need_ < reconstruct_clients_num_need_ || + reconstruct_clients_num_need_ <= secrets_minnums_ || client_num_need_ < share_clients_num_need_) { + MS_LOG(ERROR) << "CIPHER Init Params are illegal."; + return false; + } + return true; +} + +} // namespace armour +} // namespace mindspore diff --git a/mindspore/ccsrc/armour/cipher/cipher_init.h b/mindspore/ccsrc/armour/cipher/cipher_init.h new file mode 100644 index 00000000000..c4babe2e8d4 --- /dev/null +++ b/mindspore/ccsrc/armour/cipher/cipher_init.h @@ -0,0 +1,77 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef MINDSPORE_CIPHER_INIT_H +#define MINDSPORE_CIPHER_INIT_H + +#include +#include +#include "armour/secure_protocol/secret_sharing.h" +#include "proto/ps.pb.h" +#include "utils/log_adapter.h" +#include "schema/fl_job_generated.h" +#include "schema/cipher_generated.h" +#include "armour/cipher/cipher_meta_storage.h" + +namespace mindspore { +namespace armour { + +template +bool CreateArray(std::vector *newData, const flatbuffers::Vector &fbs_arr) { + size_t size = newData->size(); + size_t size_fbs_arr = fbs_arr.size(); + if (size != size_fbs_arr) return false; + for (size_t i = 0; i < size; ++i) { + newData->at(i) = fbs_arr.Get(i); + } + return true; +} + +// Initialization of secure aggregation. +class CipherInit { + public: + static CipherInit &GetInstance() { + static CipherInit instance; + return instance; + } + + // Initialize the parameters of the secure aggregation. + bool Init(const CipherPublicPara ¶m, size_t time_out_mutex, size_t cipher_initial_client_cnt, + size_t cipher_exchange_secrets_cnt, size_t cipher_share_secrets_cnt, size_t cipher_get_clientlist_cnt, + size_t cipher_reconstruct_secrets_down_cnt, size_t cipher_reconstruct_secrets_up_cnt); + + // Check whether the parameters are valid. + bool Check_Parames(); + + // Get public params. which is given to start fl job thread. + CipherPublicPara *GetPublicParams() { return &publicparam_; } + + size_t share_clients_num_need_; // the minimum number of clients to share secrets. + size_t reconstruct_clients_num_need_; // the minimum number of clients to reconstruct secret mask. + size_t client_num_need_; // the minimum number of clients to update model. + size_t get_model_num_need_; // the minimum number of clients to get model. + + size_t secrets_minnums_; // the minimum number of secret fragment s to reconstruct secret mask. + size_t featuremap_; // the size of data to deal. + size_t time_out_mutex_; // timeout mutex. + + CipherPublicPara publicparam_; // the param containing encrypted public parameters. + CipherMetaStorage cipher_meta_storage_; +}; +} // namespace armour +} // namespace mindspore + +#endif // MINDSPORE_CIPHER_COMMON_H diff --git a/mindspore/ccsrc/armour/cipher/cipher_keys.cc b/mindspore/ccsrc/armour/cipher/cipher_keys.cc new file mode 100644 index 00000000000..c9c633ca474 --- /dev/null +++ b/mindspore/ccsrc/armour/cipher/cipher_keys.cc @@ -0,0 +1,214 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "armour/cipher/cipher_keys.h" +#include "armour/cipher/cipher_meta_storage.h" + +namespace mindspore { +namespace armour { +bool CipherKeys::GetKeys(const int cur_iterator, const std::string &next_req_time, + const schema::GetExchangeKeys *get_exchange_keys_req, + std::shared_ptr get_exchange_keys_resp_builder) { + MS_LOG(INFO) << "CipherMgr::GetKeys START"; + if (get_exchange_keys_req == nullptr || get_exchange_keys_resp_builder == nullptr) { + MS_LOG(ERROR) << "Request is nullptr or Response builder is nullptr."; + BuildGetKeys(get_exchange_keys_resp_builder, schema::ResponseCode_SystemError, cur_iterator, next_req_time, false); + return false; + } + + // get clientlist from memory server. + std::vector clients; + + cipher_init_->cipher_meta_storage_.GetClientListFromServer(ps::server::kCtxExChangeKeysClientList, &clients); + + size_t cur_clients_num = clients.size(); + std::string fl_id = get_exchange_keys_req->fl_id()->str(); + + if (find(clients.begin(), clients.end(), fl_id) == clients.end()) { + BuildGetKeys(get_exchange_keys_resp_builder, schema::ResponseCode_RequestError, cur_iterator, next_req_time, false); + MS_LOG(INFO) << "The fl_id is not in clients."; + return false; + } + if (cur_clients_num < cipher_init_->client_num_need_) { + BuildGetKeys(get_exchange_keys_resp_builder, schema::ResponseCode_SucNotReady, cur_iterator, next_req_time, false); + MS_LOG(INFO) << "The server is not ready yet: cur_clients_num < client_num_need"; + MS_LOG(INFO) << "cur_clients_num : " << cur_clients_num << "cur_clients_num : " << cipher_init_->client_num_need_; + return false; + } + + MS_LOG(INFO) << "GetKeys client list: "; + for (size_t i = 0; i < clients.size(); i++) { + MS_LOG(INFO) << "fl_id: " << clients[i]; + } + + bool flag = + BuildGetKeys(get_exchange_keys_resp_builder, schema::ResponseCode_OutOfTime, cur_iterator, next_req_time, true); + return flag; +} // namespace armour + +bool CipherKeys::ExchangeKeys(const int cur_iterator, const std::string &next_req_time, + const schema::RequestExchangeKeys *exchange_keys_req, + std::shared_ptr exchange_keys_resp_builder) { + MS_LOG(INFO) << "CipherMgr::ExchangeKeys START"; + // step 0: judge if the input param is legal. + if (exchange_keys_req == nullptr || exchange_keys_resp_builder == nullptr) { + MS_LOG(ERROR) << "Request is nullptr or Response builder is nullptr."; + std::string reason = "Request is nullptr or Response builder is nullptr."; + BuildExchangeKeysRsp(exchange_keys_resp_builder, schema::ResponseCode_RequestError, reason, next_req_time, + cur_iterator); + return false; + } + + // step 1: get clientlist and client keys from memory server. + std::map>> record_public_keys; + std::vector client_list; + cipher_init_->cipher_meta_storage_.GetClientListFromServer(ps::server::kCtxExChangeKeysClientList, &client_list); + cipher_init_->cipher_meta_storage_.GetClientKeysFromServer(ps::server::kCtxClientsKeys, &record_public_keys); + + // step2: process new item data. and update new item data to memory server. + size_t cur_clients_num = client_list.size(); + size_t cur_clients_has_keys_num = record_public_keys.size(); + if (cur_clients_num != cur_clients_has_keys_num) { + std::string reason = "client num and keys num are not equal."; + MS_LOG(ERROR) << reason; + MS_LOG(ERROR) << "cur_clients_num is " << cur_clients_num << ". cur_clients_has_keys_num is " + << cur_clients_has_keys_num; + BuildExchangeKeysRsp(exchange_keys_resp_builder, schema::ResponseCode_OutOfTime, reason, next_req_time, + cur_iterator); + + return false; + } + + MS_LOG(INFO) << "client_num_need_ " << cipher_init_->client_num_need_ << ". cur_clients_num " << cur_clients_num; + std::string fl_id = exchange_keys_req->fl_id()->str(); + if (cur_clients_num >= cipher_init_->client_num_need_) { // the client num is enough, return false. + BuildExchangeKeysRsp(exchange_keys_resp_builder, schema::ResponseCode_OutOfTime, + "The server has received enough requests and refuse this request.", next_req_time, + cur_iterator); + MS_LOG(ERROR) << "The server has received enough requests and refuse this request."; + return false; + } + if (record_public_keys.find(fl_id) != record_public_keys.end()) { // the client already exists, return false. + BuildExchangeKeysRsp(exchange_keys_resp_builder, schema::ResponseCode_SUCCEED, + "The server has received the request, please do not request again.", next_req_time, + cur_iterator); + MS_LOG(INFO) << "The server has received the request, please do not request again."; + return false; + } + + // Gets the members of the deserialized data : exchange_keys_req + auto fbs_cpk = exchange_keys_req->c_pk(); + size_t cpk_len = fbs_cpk->size(); + auto fbs_spk = exchange_keys_req->s_pk(); + size_t spk_len = fbs_spk->size(); + + // transform fbs (fbs_cpk & fbs_spk) to a vector: public_key + std::vector> cur_public_key; + std::vector cpk(cpk_len); + std::vector spk(spk_len); + bool ret_create_code_cpk = CreateArray(&cpk, *fbs_cpk); + bool ret_create_code_spk = CreateArray(&spk, *fbs_spk); + if (!(ret_create_code_cpk && ret_create_code_spk)) { + MS_LOG(ERROR) << "create cur_public_key failed"; + BuildExchangeKeysRsp(exchange_keys_resp_builder, schema::ResponseCode_OutOfTime, "update key or client failed", + next_req_time, cur_iterator); + return false; + } + cur_public_key.push_back(cpk); + cur_public_key.push_back(spk); + + bool retcode_key = + cipher_init_->cipher_meta_storage_.UpdateClientKeyToServer(ps::server::kCtxClientsKeys, fl_id, cur_public_key); + bool retcode_client = + cipher_init_->cipher_meta_storage_.UpdateClientToServer(ps::server::kCtxExChangeKeysClientList, fl_id); + if (retcode_key && retcode_client) { + BuildExchangeKeysRsp(exchange_keys_resp_builder, schema::ResponseCode_SUCCEED, + "Success, but the server is not ready yet.", next_req_time, cur_iterator); + MS_LOG(INFO) << "The client " << fl_id << " CipherMgr::ExchangeKeys Success"; + return true; + } else { + MS_LOG(ERROR) << "update key or client failed"; + BuildExchangeKeysRsp(exchange_keys_resp_builder, schema::ResponseCode_OutOfTime, "update key or client failed", + next_req_time, cur_iterator); + return false; + } +} + +void CipherKeys::BuildExchangeKeysRsp(std::shared_ptr exchange_keys_resp_builder, + const schema::ResponseCode retcode, const std::string &reason, + const std::string &next_req_time, const int iteration) { + auto rsp_reason = exchange_keys_resp_builder->CreateString(reason); + auto rsp_next_req_time = exchange_keys_resp_builder->CreateString(next_req_time); + schema::ResponseExchangeKeysBuilder rsp_builder(*(exchange_keys_resp_builder.get())); + rsp_builder.add_retcode(retcode); + rsp_builder.add_reason(rsp_reason); + rsp_builder.add_next_req_time(rsp_next_req_time); + rsp_builder.add_iteration(iteration); + auto rsp_exchange_keys = rsp_builder.Finish(); + exchange_keys_resp_builder->Finish(rsp_exchange_keys); + return; +} + +bool CipherKeys::BuildGetKeys(std::shared_ptr fbb, const schema::ResponseCode retcode, + const int iteration, const std::string &next_req_time, bool is_good) { + schema::ReturnExchangeKeysBuilder rsp_buider(*(fbb.get())); + bool flag = true; + if (is_good) { + // convert client keys to standard keys list. + std::vector> public_keys_list; + MS_LOG(INFO) << "Get Keys: "; + std::map>> record_public_keys; + cipher_init_->cipher_meta_storage_.GetClientKeysFromServer(ps::server::kCtxClientsKeys, &record_public_keys); + if (record_public_keys.size() < cipher_init_->client_num_need_) { + MS_LOG(INFO) << "NOT READY. keys num: " << record_public_keys.size() + << "clients num: " << cipher_init_->client_num_need_; + flag = false; + } else { + for (auto iter = record_public_keys.begin(); iter != record_public_keys.end(); ++iter) { + // read (fl_id, c_pk, s_pk) from the map: record_public_keys_ + std::string fl_id = iter->first; + MS_LOG(INFO) << "fl_id : " << fl_id; + + // To serialize the members to a new Table:ClientPublicKeys + auto fbs_fl_id = fbb->CreateString(fl_id); + auto fbs_c_pk = fbb->CreateVector(iter->second[0].data(), iter->second[0].size()); + auto fbs_s_pk = fbb->CreateVector(iter->second[1].data(), iter->second[1].size()); + auto cur_public_key = schema::CreateClientPublicKeys(*fbb, fbs_fl_id, fbs_c_pk, fbs_s_pk); + public_keys_list.push_back(cur_public_key); + } + auto remote_publickeys = fbb->CreateVector(public_keys_list); + rsp_buider.add_remote_publickeys(remote_publickeys); + MS_LOG(INFO) << "CipherMgr::GetKeys Success"; + flag = true; + } + } + + auto fbs_next_req_time = fbb->CreateString(next_req_time); + rsp_buider.add_retcode(retcode); + rsp_buider.add_iteration(iteration); + rsp_buider.add_next_req_time(fbs_next_req_time); + auto rsp_get_keys = rsp_buider.Finish(); + fbb->Finish(rsp_get_keys); + return flag; +} + +void CipherKeys::ClearKeys() { + ps::server::DistributedMetadataStore::GetInstance().ResetMetadata(ps::server::kCtxExChangeKeysClientList); + ps::server::DistributedMetadataStore::GetInstance().ResetMetadata(ps::server::kCtxClientsKeys); +} + +} // namespace armour +} // namespace mindspore diff --git a/mindspore/ccsrc/armour/cipher/cipher_keys.h b/mindspore/ccsrc/armour/cipher/cipher_keys.h new file mode 100644 index 00000000000..2e2bb7b30c8 --- /dev/null +++ b/mindspore/ccsrc/armour/cipher/cipher_keys.h @@ -0,0 +1,71 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef MINDSPORE_CIPHER_KEYS_H +#define MINDSPORE_CIPHER_KEYS_H + +#include +#include +#include +#include +#include "armour/secure_protocol/secret_sharing.h" +#include "proto/ps.pb.h" +#include "utils/log_adapter.h" +#include "armour/cipher/cipher_init.h" +#include "armour/cipher/cipher_meta_storage.h" +#include "ps/server/common.h" + +namespace mindspore { +namespace armour { + +// The process of exchange keys and get keys in the secure aggregation +class CipherKeys { + public: + // initialize: get cipher_init_ + CipherKeys() { cipher_init_ = &CipherInit::GetInstance(); } + + static CipherKeys &GetInstance() { + static CipherKeys instance; + return instance; + } + + // handle the client's request of get keys. + bool GetKeys(const int cur_iterator, const std::string &next_req_time, + const schema::GetExchangeKeys *get_exchange_keys_req, + std::shared_ptr get_exchange_keys_resp_builder); + + // handle the client's request of exchange keys. + bool ExchangeKeys(const int cur_iterator, const std::string &next_req_time, + const schema::RequestExchangeKeys *exchange_keys_req, + std::shared_ptr exchange_keys_resp_builder); + + // build response code of get keys. + bool BuildGetKeys(std::shared_ptr fbb, const schema::ResponseCode retcode, const int iteration, + const std::string &next_req_time, bool is_good); + // build response code of exchange keys. + void BuildExchangeKeysRsp(std::shared_ptr exchange_keys_resp_builder, + const schema::ResponseCode retcode, const std::string &reason, + const std::string &next_req_time, const int iteration); + // clear the shared memory. + void ClearKeys(); + + private: + CipherInit *cipher_init_; // the parameter of the secure aggregation +}; +} // namespace armour +} // namespace mindspore + +#endif // MINDSPORE_CIPHER_KEYS_H diff --git a/mindspore/ccsrc/armour/cipher/cipher_meta_storage.cc b/mindspore/ccsrc/armour/cipher/cipher_meta_storage.cc new file mode 100644 index 00000000000..cf88c7eff6b --- /dev/null +++ b/mindspore/ccsrc/armour/cipher/cipher_meta_storage.cc @@ -0,0 +1,188 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "armour/cipher/cipher_meta_storage.h" + +namespace mindspore { +namespace armour { + +void CipherMetaStorage::GetClientSharesFromServer( + const char *list_name, std::map> *clients_shares_list) { + const ps::PBMetadata &clients_shares_pb_out = + ps::server::DistributedMetadataStore::GetInstance().GetMetadata(list_name); + const ps::ClientShares &clients_shares_pb = clients_shares_pb_out.client_shares(); + auto iter = clients_shares_pb.client_secret_shares().begin(); + for (; iter != clients_shares_pb.client_secret_shares().end(); ++iter) { + std::string fl_id = iter->first; + const ps::SharesPb &shares_pb = iter->second; + std::vector encrpted_shares_new; + for (int index_shares = 0; index_shares < shares_pb.clientsharestrs_size(); ++index_shares) { + const ps::ClientShareStr &client_share_str_pb = shares_pb.clientsharestrs(index_shares); + clientshare_str new_clientshare; + new_clientshare.fl_id = client_share_str_pb.fl_id(); + new_clientshare.index = client_share_str_pb.index(); + new_clientshare.share.assign(client_share_str_pb.share().begin(), client_share_str_pb.share().end()); + encrpted_shares_new.push_back(new_clientshare); + } + clients_shares_list->insert(std::pair>(fl_id, encrpted_shares_new)); + } +} + +void CipherMetaStorage::GetClientListFromServer(const char *list_name, std::vector *clients_list) { + const ps::PBMetadata &client_list_pb_out = ps::server::DistributedMetadataStore::GetInstance().GetMetadata(list_name); + const ps::UpdateModelClientList &client_list_pb = client_list_pb_out.client_list(); + for (int i = 0; i < client_list_pb.fl_id_size(); ++i) { + std::string fl_id = client_list_pb.fl_id(i); + clients_list->push_back(fl_id); + } +} + +void CipherMetaStorage::GetClientKeysFromServer( + const char *list_name, std::map>> *clients_keys_list) { + const ps::PBMetadata &clients_keys_pb_out = + ps::server::DistributedMetadataStore::GetInstance().GetMetadata(list_name); + const ps::ClientKeys &clients_keys_pb = clients_keys_pb_out.client_keys(); + + for (auto iter = clients_keys_pb.client_keys().begin(); iter != clients_keys_pb.client_keys().end(); ++iter) { + // const PairClientKeys & pair_client_keys_pb = clients_keys_pb.client_keys(i); + std::string fl_id = iter->first; + ps::KeysPb keys_pb = iter->second; + std::vector cpk(keys_pb.key(0).begin(), keys_pb.key(0).end()); + std::vector spk(keys_pb.key(1).begin(), keys_pb.key(1).end()); + std::vector> cur_keys; + cur_keys.push_back(cpk); + cur_keys.push_back(spk); + clients_keys_list->insert(std::pair>>(fl_id, cur_keys)); + } +} + +bool CipherMetaStorage::GetClientNoisesFromServer(const char *list_name, std::vector *cur_public_noise) { + const ps::PBMetadata &clients_noises_pb_out = + ps::server::DistributedMetadataStore::GetInstance().GetMetadata(list_name); + const ps::ClientNoises &clients_noises_pb = clients_noises_pb_out.client_noises(); + while (clients_noises_pb.has_one_client_noises() == false) { + MS_LOG(INFO) << "GetClientNoisesFromServer NULL."; + std::this_thread::sleep_for(std::chrono::milliseconds(50)); + } + cur_public_noise->assign(clients_noises_pb.one_client_noises().noise().begin(), + clients_noises_pb.one_client_noises().noise().end()); + return true; +} + +bool CipherMetaStorage::GetPrimeFromServer(const char *list_name, unsigned char *prime) { + const ps::PBMetadata &prime_pb_out = ps::server::DistributedMetadataStore::GetInstance().GetMetadata(list_name); + auto &prime_list_pb = prime_pb_out.prime_list(); + if (prime_list_pb.prime_size() > 0 && prime_list_pb.prime(0).size() >= PRIME_MAX_LEN) { + for (int i = 0; i < PRIME_MAX_LEN; i++) { + prime[i] = static_cast(prime_list_pb.prime(0)[i]); + } + return true; + } + return false; +} + +bool CipherMetaStorage::UpdateClientToServer(const char *list_name, const std::string &fl_id) { + bool retcode = true; + ps::FLId fl_id_pb; + fl_id_pb.set_fl_id(fl_id); + ps::PBMetadata client_pb; + client_pb.mutable_fl_id()->MergeFrom(fl_id_pb); + retcode = ps::server::DistributedMetadataStore::GetInstance().UpdateMetadata(list_name, client_pb); + return retcode; +} +void CipherMetaStorage::RegisterPrime(const char *list_name, const std::string &prime) { + ps::Prime prime_id_pb; + prime_id_pb.set_prime(prime); + ps::PBMetadata prime_pb; + prime_pb.mutable_prime()->MergeFrom(prime_id_pb); + ps::server::DistributedMetadataStore::GetInstance().RegisterMetadata(list_name, prime_pb); +} + +bool CipherMetaStorage::UpdateClientKeyToServer(const char *list_name, const std::string &fl_id, + const std::vector> &cur_public_key) { + bool retcode = true; + if (cur_public_key.size() < 2) { + MS_LOG(ERROR) << "cur_public_key's size must is 2. actual size is " << cur_public_key.size(); + return false; + } + // update new item to memory server. + ps::KeysPb keys; + keys.add_key()->assign(cur_public_key[0].begin(), cur_public_key[0].end()); + keys.add_key()->assign(cur_public_key[1].begin(), cur_public_key[1].end()); + ps::PairClientKeys pair_client_keys_pb; + pair_client_keys_pb.set_fl_id(fl_id); + pair_client_keys_pb.mutable_client_keys()->MergeFrom(keys); + ps::PBMetadata client_and_keys_pb; + client_and_keys_pb.mutable_pair_client_keys()->MergeFrom(pair_client_keys_pb); + retcode = ps::server::DistributedMetadataStore::GetInstance().UpdateMetadata(list_name, client_and_keys_pb); + return retcode; +} + +bool CipherMetaStorage::UpdateClientNoiseToServer(const char *list_name, const std::vector &cur_public_noise) { + // update new item to memory server. + ps::OneClientNoises noises_pb; + *noises_pb.mutable_noise() = {cur_public_noise.begin(), cur_public_noise.end()}; + ps::PBMetadata client_noises_pb; + client_noises_pb.mutable_one_client_noises()->MergeFrom(noises_pb); + return ps::server::DistributedMetadataStore::GetInstance().UpdateMetadata(list_name, client_noises_pb); +} + +bool CipherMetaStorage::UpdateClientShareToServer( + const char *list_name, const std::string &fl_id, + const flatbuffers::Vector> *shares) { + bool retcode = true; + int size_shares = shares->size(); + ps::SharesPb shares_pb; + for (int index = 0; index < size_shares; ++index) { + // new item + ps::ClientShareStr *client_share_str_new_p = shares_pb.add_clientsharestrs(); + std::string fl_id_new = (*shares)[index]->fl_id()->str(); + int index_new = (*shares)[index]->index(); + auto share = (*shares)[index]->share(); + client_share_str_new_p->set_share(reinterpret_cast(share->data()), share->size()); + client_share_str_new_p->set_fl_id(fl_id_new); + client_share_str_new_p->set_index(index_new); + } + ps::PairClientShares pair_client_shares_pb; + pair_client_shares_pb.set_fl_id(fl_id); + pair_client_shares_pb.mutable_client_shares()->MergeFrom(shares_pb); + ps::PBMetadata client_and_shares_pb; + client_and_shares_pb.mutable_pair_client_shares()->MergeFrom(pair_client_shares_pb); + retcode = ps::server::DistributedMetadataStore::GetInstance().UpdateMetadata(list_name, client_and_shares_pb); + return retcode; +} + +void CipherMetaStorage::RegisterClass() { + ps::PBMetadata exchange_kyes_client_list; + ps::server::DistributedMetadataStore::GetInstance().RegisterMetadata(ps::server::kCtxExChangeKeysClientList, + exchange_kyes_client_list); + ps::PBMetadata clients_keys; + ps::server::DistributedMetadataStore::GetInstance().RegisterMetadata(ps::server::kCtxClientsKeys, clients_keys); + ps::PBMetadata reconstruct_client_list; + ps::server::DistributedMetadataStore::GetInstance().RegisterMetadata(ps::server::kCtxReconstructClientList, + reconstruct_client_list); + ps::PBMetadata clients_reconstruct_shares; + ps::server::DistributedMetadataStore::GetInstance().RegisterMetadata(ps::server::kCtxClientsReconstructShares, + clients_reconstruct_shares); + ps::PBMetadata share_secretes_client_list; + ps::server::DistributedMetadataStore::GetInstance().RegisterMetadata(ps::server::kCtxShareSecretsClientList, + share_secretes_client_list); + ps::PBMetadata clients_encrypt_shares; + ps::server::DistributedMetadataStore::GetInstance().RegisterMetadata(ps::server::kCtxClientsEncryptedShares, + clients_encrypt_shares); +} +} // namespace armour +} // namespace mindspore diff --git a/mindspore/ccsrc/armour/cipher/cipher_meta_storage.h b/mindspore/ccsrc/armour/cipher/cipher_meta_storage.h new file mode 100644 index 00000000000..acf3b58c629 --- /dev/null +++ b/mindspore/ccsrc/armour/cipher/cipher_meta_storage.h @@ -0,0 +1,92 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef MINDSPORE_CIPHER_META_STORAGE_H +#define MINDSPORE_CIPHER_META_STORAGE_H + +#include +#include +#include +#include +#include +#include +#include +#include "proto/ps.pb.h" +#include "utils/log_adapter.h" +#include "armour/secure_protocol/secret_sharing.h" +#include "schema/fl_job_generated.h" +#include "schema/cipher_generated.h" +#include "ps/server/distributed_metadata_store.h" +#include "ps/server/common.h" + +namespace mindspore { +namespace armour { + +constexpr int SHARE_MAX_SIZE = 256; +constexpr int SECRET_MAX_LEN_DOUBLE = 66; + +struct clientshare_str { + std::string fl_id; + std::vector share; + int index; +}; + +struct CipherPublicPara { + int t; + int g; + unsigned char prime[PRIME_MAX_LEN]; + unsigned char p[SECRET_MAX_LEN]; + float dp_eps; + float dp_delta; + float dp_norm_clip; + int encrypt_type; +}; + +class CipherMetaStorage { + public: + // Register the shared value involved in the security aggregation. + void RegisterClass(); + + // Register Prime. + void RegisterPrime(const char *list_name, const std::string &prime); + // Get tprime from shared server. + bool GetPrimeFromServer(const char *list_name, unsigned char *prime); + // Get client shares from shared server. + void GetClientSharesFromServer(const char *list_name, + std::map> *clients_shares_list); + // Get client list from shared server. + void GetClientListFromServer(const char *list_name, std::vector *clients_list); + // Get client keys from shared server. + void GetClientKeysFromServer(const char *list_name, + std::map>> *clients_keys_list); + // Get client noises from shared server. + bool GetClientNoisesFromServer(const char *list_name, std::vector *cur_public_noise); + // Update client fl_id to shared server. + bool UpdateClientToServer(const char *list_name, const std::string &fl_id); + // Update client key to shared server. + bool UpdateClientKeyToServer(const char *list_name, const std::string &fl_id, + const std::vector> &cur_public_key); + // Update client noise to shared server. + bool UpdateClientNoiseToServer(const char *list_name, const std::vector &cur_public_noise); + // Update client share to shared server. + bool UpdateClientShareToServer( + const char *list_name, const std::string &fl_id, + const flatbuffers::Vector> *shares); +}; +} // namespace armour +} // namespace mindspore + +#endif // MINDSPORE_CIPHER_META_STORAGE_H diff --git a/mindspore/ccsrc/armour/cipher/cipher_reconstruct.cc b/mindspore/ccsrc/armour/cipher/cipher_reconstruct.cc new file mode 100644 index 00000000000..9771797c6c8 --- /dev/null +++ b/mindspore/ccsrc/armour/cipher/cipher_reconstruct.cc @@ -0,0 +1,418 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "armour/cipher/cipher_reconstruct.h" +#include "ps/server/common.h" +#include "armour/secure_protocol/random.h" +#include "armour/secure_protocol/key_agreement.h" +#include "armour/cipher/cipher_meta_storage.h" + +namespace mindspore { +namespace armour { +bool CipherReconStruct::CombineMask( + std::vector *shares_tmp, std::map> *client_keys, + const std::vector &clients_share_list, + const std::map>> &record_public_keys, + const std::map> &reconstruct_secret_list, + const std::vector &client_list) { + bool retcode = true; + for (auto iter = reconstruct_secret_list.begin(); iter != reconstruct_secret_list.end(); ++iter) { + // define flag_share: judge we need b or s + bool flag_share = true; + const std::string fl_id = iter->first; + std::vector::const_iterator ptr = client_list.begin(); + for (; ptr < client_list.end(); ++ptr) { + if (*ptr == fl_id) { + flag_share = false; + break; + } + } + MS_LOG(INFO) << "fl_id_src : " << fl_id; + mpz_t prime; + mpz_init(prime); + auto publicparam_ = CipherInit::GetInstance().GetPublicParams(); + mpz_import(prime, PRIME_MAX_LEN, 1, 1, 0, 0, publicparam_->prime); + + if (iter->second.size() >= cipher_init_->secrets_minnums_) { // combine private key seed. + MS_LOG(INFO) << "start assign secrets shares to public shares "; + for (int i = 0; i < static_cast(cipher_init_->secrets_minnums_); ++i) { + shares_tmp->at(i)->index = (iter->second)[i].index; + shares_tmp->at(i)->len = (iter->second)[i].share.size(); + if (memcpy_s(shares_tmp->at(i)->data, shares_tmp->at(i)->len, (iter->second)[i].share.data(), + shares_tmp->at(i)->len) != 0) { + MS_LOG(ERROR) << "shares_tmp copy failed"; + retcode = false; + } + MS_LOG(INFO) << "fl_id_des : " << (iter->second)[i].fl_id; + std::string print_share_data(reinterpret_cast(shares_tmp->at(i)->data), shares_tmp->at(i)->len); + } + MS_LOG(INFO) << "end assign secrets shares to public shares "; + + size_t length; + char secret[SECRET_MAX_LEN] = {0}; + SecretSharing combine(prime); + if (combine.Combine(static_cast(cipher_init_->secrets_minnums_), *shares_tmp, secret, &length) < 0) + retcode = false; + length = SECRET_MAX_LEN; + MS_LOG(INFO) << "combine secrets shares Success."; + + if (flag_share) { + MS_LOG(INFO) << "start get complete s_uv."; + std::vector noise(cipher_init_->featuremap_, 0.0); + if (GetSuvNoise(clients_share_list, record_public_keys, fl_id, &noise, secret, length) == false) + retcode = false; + client_keys->at(fl_id) = noise; + MS_LOG(INFO) << " fl_id : " << fl_id; + MS_LOG(INFO) << "end get complete s_uv."; + } else { + std::vector noise; + if (Random::RandomAESCTR(&noise, cipher_init_->featuremap_, (const unsigned char *)secret, SECRET_MAX_LEN) < 0) + retcode = false; + for (size_t index_noise = 0; index_noise < cipher_init_->featuremap_; index_noise++) { + noise[index_noise] *= -1; + } + client_keys->at(fl_id) = noise; + MS_LOG(INFO) << " fl_id : " << fl_id; + } + } + } + + return retcode; +} + +bool CipherReconStruct::ReconstructSecretsGenNoise(const std::vector &client_list) { + // get reconstruct_secret_list_ori from memory server + MS_LOG(INFO) << "CipherReconStruct::ReconstructSecretsGenNoise START"; + bool retcode = true; + std::map> reconstruct_secret_list_ori; + cipher_init_->cipher_meta_storage_.GetClientSharesFromServer(ps::server::kCtxClientsReconstructShares, + &reconstruct_secret_list_ori); + std::map>> record_public_keys; + cipher_init_->cipher_meta_storage_.GetClientKeysFromServer(ps::server::kCtxClientsKeys, &record_public_keys); + std::vector clients_reconstruct_list; + cipher_init_->cipher_meta_storage_.GetClientListFromServer(ps::server::kCtxReconstructClientList, + &clients_reconstruct_list); + std::vector clients_share_list; + cipher_init_->cipher_meta_storage_.GetClientListFromServer(ps::server::kCtxShareSecretsClientList, + &clients_share_list); + if (reconstruct_secret_list_ori.size() != clients_reconstruct_list.size() || + record_public_keys.size() < cipher_init_->client_num_need_ || + clients_share_list.size() < cipher_init_->share_clients_num_need_) { + MS_LOG(ERROR) << "get data from server memory failed"; + return false; + } + + std::map> reconstruct_secret_list; + ConvertSharesToShares(reconstruct_secret_list_ori, &reconstruct_secret_list); + std::vector shares_tmp; + if (MallocShares(&shares_tmp, cipher_init_->secrets_minnums_) == false) { + MS_LOG(ERROR) << "Reconstruct malloc shares_tmp invalid."; + return false; + } + MS_LOG(INFO) << "Reconstruct client list: "; + std::vector::const_iterator ptr_tmp = client_list.begin(); + for (; ptr_tmp < client_list.end(); ++ptr_tmp) { + MS_LOG(INFO) << *ptr_tmp; + } + MS_LOG(INFO) << "Reconstruct secrets shares: "; + std::map> client_keys; + + retcode = CombineMask(&shares_tmp, &client_keys, clients_share_list, record_public_keys, reconstruct_secret_list, + client_list); + + DeleteShares(&shares_tmp); + if (retcode) { + std::vector noise; + if (GetNoiseMasksSum(&noise, client_keys) == false) { + MS_LOG(ERROR) << " GetNoiseMasksSum failed"; + return false; + } + client_keys.clear(); + MS_LOG(INFO) << " ReconstructSecretsGenNoise updata noise to server"; + + if (cipher_init_->cipher_meta_storage_.UpdateClientNoiseToServer(ps::server::kCtxClientNoises, noise) == false) + return false; + MS_LOG(INFO) << " ReconstructSecretsGenNoise Success"; + } else { + MS_LOG(INFO) << " ReconstructSecretsGenNoise failed. because gen noise inside failed"; + } + + return retcode; +} + +// reconstruct secrets +bool CipherReconStruct::ReconstructSecrets(const int cur_iterator, const std::string &next_req_time, + const schema::SendReconstructSecret *reconstruct_secret_req, + std::shared_ptr reconstruct_secret_resp_builder, + const std::vector &client_list) { + MS_LOG(INFO) << "CipherReconStruct::ReconstructSecrets START"; + clock_t start_time = clock(); + if (reconstruct_secret_req == nullptr || reconstruct_secret_resp_builder == nullptr) { + MS_LOG(ERROR) << "Request is nullptr or Response builder is nullptr. "; + BuildReconstructSecretsRsp(reconstruct_secret_resp_builder, schema::ResponseCode_RequestError, + "Request is nullptr or Response builder is nullptr.", cur_iterator, next_req_time); + return false; + } + if (client_list.size() < cipher_init_->reconstruct_clients_num_need_) { + MS_LOG(ERROR) << "illegal parameters. update model client_list size: " << client_list.size(); + BuildReconstructSecretsRsp( + reconstruct_secret_resp_builder, schema::ResponseCode_RequestError, + "illegal parameters: update model client_list size must larger than reconstruct_clients_num_need", cur_iterator, + next_req_time); + return false; + } + std::vector clients_reconstruct_list; + cipher_init_->cipher_meta_storage_.GetClientListFromServer(ps::server::kCtxReconstructClientList, + &clients_reconstruct_list); + std::map> clients_shares_all; + cipher_init_->cipher_meta_storage_.GetClientSharesFromServer(ps::server::kCtxClientsReconstructShares, + &clients_shares_all); + + size_t count_client_num = clients_shares_all.size(); + if (count_client_num != clients_reconstruct_list.size()) { + BuildReconstructSecretsRsp(reconstruct_secret_resp_builder, schema::ResponseCode_OutOfTime, + "shares client size and client size are not equal.", cur_iterator, next_req_time); + MS_LOG(ERROR) << "shares client size and client size are not equal."; + return false; + } + int iterator = reconstruct_secret_req->iteration(); + std::string fl_id = reconstruct_secret_req->fl_id()->str(); + if (iterator != cur_iterator) { + BuildReconstructSecretsRsp(reconstruct_secret_resp_builder, schema::ResponseCode_OutOfTime, + "The iteration round of the client does not match the current iteration.", cur_iterator, + next_req_time); + MS_LOG(ERROR) << "Client " << fl_id << " The iteration round of the client does not match the current iteration."; + return false; + } + + if (find(client_list.begin(), client_list.end(), fl_id) == client_list.end()) { // client not in client list. + BuildReconstructSecretsRsp(reconstruct_secret_resp_builder, schema::ResponseCode_OutOfTime, + "The client is not in update model client list.", cur_iterator, next_req_time); + MS_LOG(ERROR) << "The client " << fl_id << " is not in update model client list."; + return false; + } + if (find(clients_reconstruct_list.begin(), clients_reconstruct_list.end(), fl_id) != clients_reconstruct_list.end()) { + BuildReconstructSecretsRsp(reconstruct_secret_resp_builder, schema::ResponseCode_SUCCEED, + "Client has sended messages.", cur_iterator, next_req_time); + MS_LOG(INFO) << "Error, client " << fl_id << " has sended messages."; + return false; + } + auto reconstruct_secret_shares = reconstruct_secret_req->reconstruct_secret_shares(); + bool retcode_client = + cipher_init_->cipher_meta_storage_.UpdateClientToServer(ps::server::kCtxReconstructClientList, fl_id); + bool retcode_share = cipher_init_->cipher_meta_storage_.UpdateClientShareToServer( + ps::server::kCtxClientsReconstructShares, fl_id, reconstruct_secret_shares); + if (!(retcode_client && retcode_share)) { + BuildReconstructSecretsRsp(reconstruct_secret_resp_builder, schema::ResponseCode_OutOfTime, + "reconstruct update shares or client failed.", cur_iterator, next_req_time); + MS_LOG(ERROR) << "reconstruct update shares or client failed."; + return false; + } + count_client_num = count_client_num + 1; + if (count_client_num < cipher_init_->reconstruct_clients_num_need_) { + BuildReconstructSecretsRsp(reconstruct_secret_resp_builder, schema::ResponseCode_SUCCEED, + "Success,but the server is not ready to reconstruct secret yet.", cur_iterator, + next_req_time); + MS_LOG(INFO) << "ReconstructSecrets" << fl_id << " Success, but count " << count_client_num << "is not enough."; + return true; + } else { + bool retcode_result = true; + const ps::PBMetadata &clients_noises_pb_out = + ps::server::DistributedMetadataStore::GetInstance().GetMetadata(ps::server::kCtxClientNoises); + const ps::ClientNoises &clients_noises_pb = clients_noises_pb_out.client_noises(); + if (clients_noises_pb.has_one_client_noises() == false) { + MS_LOG(INFO) << "Success,the secret will be reconstructed."; + retcode_result = ReconstructSecretsGenNoise(client_list); + if (retcode_result) { + BuildReconstructSecretsRsp(reconstruct_secret_resp_builder, schema::ResponseCode_SUCCEED, + "Success,the secret is reconstructing.", cur_iterator, next_req_time); + MS_LOG(INFO) << "CipherReconStruct::ReconstructSecrets" << fl_id << " Success, reconstruct ok."; + } else { + BuildReconstructSecretsRsp(reconstruct_secret_resp_builder, schema::ResponseCode_OutOfTime, + "the secret restructs failed.", cur_iterator, next_req_time); + MS_LOG(ERROR) << "the secret restructs failed."; + } + } else { + BuildReconstructSecretsRsp(reconstruct_secret_resp_builder, schema::ResponseCode_SUCCEED, + "Clients' number is full.", cur_iterator, next_req_time); + MS_LOG(INFO) << "CipherReconStruct::ReconstructSecrets" << fl_id << " Success : no need reconstruct."; + } + clock_t end_time = clock(); + double duration = static_cast((end_time - start_time) * 1.0 / CLOCKS_PER_SEC); + MS_LOG(INFO) << "Reconstruct get + gennoise data time is : " << duration; + return retcode_result; + } +} + +bool CipherReconStruct::GetNoiseMasksSum(std::vector *result, + const std::map> &client_keys) { + float sum[cipher_init_->featuremap_] = {0.0}; + for (auto iter = client_keys.begin(); iter != client_keys.end(); iter++) { + if (iter->second.size() != cipher_init_->featuremap_) { + return false; + } + for (size_t i = 0; i < cipher_init_->featuremap_; i++) { + sum[i] += iter->second[i]; + } + } + for (size_t i = 0; i < cipher_init_->featuremap_; i++) { + result->push_back(sum[i]); + } + return true; +} + +void CipherReconStruct::ClearReconstructSecrets() { + MS_LOG(INFO) << "CipherReconStruct::ClearReconstructSecrets START"; + ps::server::DistributedMetadataStore::GetInstance().ResetMetadata(ps::server::kCtxReconstructClientList); + ps::server::DistributedMetadataStore::GetInstance().ResetMetadata(ps::server::kCtxClientsReconstructShares); + ps::server::DistributedMetadataStore::GetInstance().ResetMetadata(ps::server::kCtxClientNoises); + MS_LOG(INFO) << "CipherReconStruct::ClearReconstructSecrets Success"; +} + +void CipherReconStruct::BuildReconstructSecretsRsp(std::shared_ptr fbb, + const schema::ResponseCode retcode, const std::string &reason, + const int iteration, const std::string &next_req_time) { + auto fbs_reason = fbb->CreateString(reason); + auto fbs_next_req_time = fbb->CreateString(next_req_time); + schema::ReconstructSecretBuilder rsp_reconstruct_secret_builder(*(fbb.get())); + rsp_reconstruct_secret_builder.add_retcode(retcode); + rsp_reconstruct_secret_builder.add_reason(fbs_reason); + rsp_reconstruct_secret_builder.add_iteration(iteration); + rsp_reconstruct_secret_builder.add_next_req_time(fbs_next_req_time); + auto rsp_reconstruct_secret = rsp_reconstruct_secret_builder.Finish(); + fbb->Finish(rsp_reconstruct_secret); + return; +} + +bool CipherReconStruct::GetSuvNoise( + const std::vector &clients_share_list, + const std::map>> &record_public_keys, const string &fl_id, + std::vector *noise, char *secret, int length) { + for (auto p_key = clients_share_list.begin(); p_key != clients_share_list.end(); ++p_key) { + if (*p_key != fl_id) { + PrivateKey *privKey1 = KeyAgreement::FromPrivateBytes((unsigned char *)secret, length); + if (privKey1 == NULL) { + MS_LOG(ERROR) << "create privKey1 failed\n"; + return false; + } + std::vector public_key = record_public_keys.at(*p_key)[1]; + PublicKey *pubKey1 = KeyAgreement::FromPublicBytes(public_key.data(), public_key.size()); + if (pubKey1 == NULL) { + MS_LOG(ERROR) << "create pubKey1 failed\n"; + return false; + } + MS_LOG(INFO) << "fl_id : " << fl_id << "other id : " << *p_key; + unsigned char secret1[SECRET_MAX_LEN] = {0}; + unsigned char salt[SECRET_MAX_LEN] = {0}; + if (KeyAgreement::ComputeSharedKey(privKey1, pubKey1, SECRET_MAX_LEN, salt, SECRET_MAX_LEN, secret1) < 0) { + MS_LOG(ERROR) << "ComputeSharedKey failed\n"; + return false; + } + + std::vector noise_tmp; + if (Random::RandomAESCTR(&noise_tmp, cipher_init_->featuremap_, (const unsigned char *)secret1, SECRET_MAX_LEN) < + 0) { + MS_LOG(ERROR) << "RandomAESCTR failed\n"; + return false; + } + bool symbol_noise = GetSymbol(fl_id, *p_key); + size_t index = 0; + size_t size_noise = noise_tmp.size(); + if (symbol_noise == false) { + for (; index < size_noise; ++index) { + noise_tmp[index] = noise_tmp[index] * -1; + noise->at(index) += noise_tmp[index]; + } + } else { + for (; index < size_noise; ++index) { + noise->at(index) += noise_tmp[index]; + } + } + for (int i = 0; i < 5; i++) { + MS_LOG(INFO) << "index " << i << " : " << noise_tmp[i]; + } + } + } + return true; +} + +bool CipherReconStruct::GetSymbol(const std::string &str1, const std::string &str2) { + if (str1 > str2) { + return true; + } else { + return false; + } +} + +void CipherReconStruct::ConvertSharesToShares(const std::map> &src, + std::map> *des) { + for (auto iter_ori = src.begin(); iter_ori != src.end(); ++iter_ori) { + std::string fl_des = iter_ori->first; + auto &cur_clientshare_str = iter_ori->second; + for (size_t index_clientshare = 0; index_clientshare < cur_clientshare_str.size(); ++index_clientshare) { + std::string fl_src = cur_clientshare_str[index_clientshare].fl_id; + clientshare_str value; + value.fl_id = fl_des; + value.share = cur_clientshare_str[index_clientshare].share; + value.index = cur_clientshare_str[index_clientshare].index; + if (des->find(fl_src) == des->end()) { // fl_id_des is not in reconstruct_secret_list_ + std::vector value_list; + value_list.push_back(value); + des->insert(std::pair>(fl_src, value_list)); + } else { // fl_id_des is in reconstruct_secret_list_ + des->at(fl_src).push_back(value); + } + } + } +} + +bool CipherReconStruct::MallocShares(std::vector *shares_tmp, int shares_size) { + for (int i = 0; i < shares_size; ++i) { + Share *share_i = new Share; + if (share_i == nullptr) { + MS_LOG(ERROR) << "shares_tmp " << i << " memory to cipher is invalid."; + DeleteShares(shares_tmp); + return false; + } + share_i->data = new unsigned char[SHARE_MAX_SIZE]; + if (share_i->data == nullptr) { + MS_LOG(ERROR) << "shares_tmp's data " << i << " memory to cipher is invalid."; + DeleteShares(shares_tmp); + return false; + } + share_i->index = 0; + share_i->len = SHARE_MAX_SIZE; + shares_tmp->push_back(share_i); + } + return true; +} + +void CipherReconStruct::DeleteShares(std::vector *shares_tmp) { + if (shares_tmp->size() != 0) { + for (size_t i = 0; i < shares_tmp->size(); ++i) { + if (shares_tmp->at(i) != nullptr && shares_tmp->at(i)->data != nullptr) { + delete[](shares_tmp->at(i)->data); + shares_tmp->at(i)->data = nullptr; + } + delete shares_tmp->at(i); + shares_tmp->at(i) = nullptr; + } + } + return; +} + +} // namespace armour +} // namespace mindspore diff --git a/mindspore/ccsrc/armour/cipher/cipher_reconstruct.h b/mindspore/ccsrc/armour/cipher/cipher_reconstruct.h new file mode 100644 index 00000000000..78146ce22bb --- /dev/null +++ b/mindspore/ccsrc/armour/cipher/cipher_reconstruct.h @@ -0,0 +1,87 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef MINDSPORE_CIPHER_RECONSTRUCT_H +#define MINDSPORE_CIPHER_RECONSTRUCT_H + +#include +#include +#include +#include +#include +#include "armour/secure_protocol/secret_sharing.h" +#include "proto/ps.pb.h" +#include "utils/log_adapter.h" +#include "armour/cipher/cipher_init.h" +#include "armour/cipher/cipher_meta_storage.h" + +namespace mindspore { +namespace armour { +// The process of reconstruct secret mask in the secure aggregation +class CipherReconStruct { + public: + // initialize: get cipher_init_ + CipherReconStruct() { cipher_init_ = &CipherInit::GetInstance(); } + + static CipherReconStruct &GetInstance() { + static CipherReconStruct instance; + return instance; + } + + // reconstruct secret mask + bool ReconstructSecrets(const int cur_iterator, const std::string &next_req_time, + const schema::SendReconstructSecret *reconstruct_secret_req, + std::shared_ptr reconstruct_secret_resp_builder, + const std::vector &client_list); + + // build response code of reconstruct secret. + void BuildReconstructSecretsRsp(std::shared_ptr fbb, const schema::ResponseCode retcode, + const std::string &reason, const int iteration, const std::string &next_req_time); + + // clear the shared memory. + void ClearReconstructSecrets(); + + private: + CipherInit *cipher_init_; // the parameter of the secure aggregation + // get mask symbol by comparing str1 and str2. + bool GetSymbol(const std::string &str1, const std::string &str2); + // get suv noise by computing shares result. + bool GetSuvNoise(const std::vector &clients_share_list, + const std::map>> &record_public_keys, + const string &fl_id, std::vector *noise, char *secret, int length); + // malloc shares. + bool MallocShares(std::vector *shares_tmp, int shares_size); + // delete shares. + void DeleteShares(std::vector *shares_tmp); + // convert shares from receiving clients to sending clients. + void ConvertSharesToShares(const std::map> &src, + std::map> *des); + // generate noise from shares. + bool ReconstructSecretsGenNoise(const std::vector &client_list); + // get noise masks sum. + bool GetNoiseMasksSum(std::vector *result, const std::map> &client_keys); + + // combine noise mask. + bool CombineMask(std::vector *shares_tmp, std::map> *client_keys, + const std::vector &clients_share_list, + const std::map>> &record_public_keys, + const std::map> &reconstruct_secret_list, + const std::vector &client_list); +}; +} // namespace armour +} // namespace mindspore + +#endif // MINDSPORE_CIPHER_KEYS_H diff --git a/mindspore/ccsrc/armour/cipher/cipher_shares.cc b/mindspore/ccsrc/armour/cipher/cipher_shares.cc new file mode 100644 index 00000000000..4caaf4c2550 --- /dev/null +++ b/mindspore/ccsrc/armour/cipher/cipher_shares.cc @@ -0,0 +1,219 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "armour/cipher/cipher_shares.h" +#include "ps/server/common.h" +#include "armour/cipher/cipher_meta_storage.h" + +namespace mindspore { +namespace armour { +bool CipherShares::ShareSecrets(const int cur_iterator, const schema::RequestShareSecrets *share_secrets_req, + std::shared_ptr share_secrets_resp_builder, + const string next_req_time) { + MS_LOG(INFO) << "CipherShares::ShareSecrets START"; + if (share_secrets_req == nullptr) { + MS_LOG(ERROR) << "Request is nullptr or Response builder is nullptr."; + std::string reason = "Request is nullptr or Response builder is nullptr."; + BuildShareSecretsRsp(share_secrets_resp_builder, schema::ResponseCode_RequestError, reason, next_req_time, + cur_iterator); + return false; + } + + // step 1: get client list and share secrets from memory server. + clock_t start_time = clock(); + std::vector clients_share_list; + cipher_init_->cipher_meta_storage_.GetClientListFromServer(ps::server::kCtxShareSecretsClientList, + &clients_share_list); + std::vector clients_exchange_list; + cipher_init_->cipher_meta_storage_.GetClientListFromServer(ps::server::kCtxExChangeKeysClientList, + &clients_exchange_list); + std::map> encrypted_shares_all; + cipher_init_->cipher_meta_storage_.GetClientSharesFromServer(ps::server::kCtxClientsEncryptedShares, + &encrypted_shares_all); + + MS_LOG(INFO) << "Client of keys size : " << clients_exchange_list.size() + << "client of shares size : " << clients_share_list.size() << "shares size" + << encrypted_shares_all.size(); + if (encrypted_shares_all.size() != clients_share_list.size()) { + BuildShareSecretsRsp(share_secrets_resp_builder, schema::ResponseCode_OutOfTime, + "client of shares and shares size are not equal", next_req_time, cur_iterator); + MS_LOG(ERROR) << "client of shares and shares size are not equal. client of shares size : " + << clients_share_list.size() << "shares size" << encrypted_shares_all.size(); + } + + // step 2: update new item to memory server. serialise: update pb struct to memory server. + int iteration = share_secrets_req->iteration(); + std::string fl_id_src = share_secrets_req->fl_id()->str(); + if (find(clients_exchange_list.begin(), clients_exchange_list.end(), fl_id_src) == + clients_exchange_list.end()) { // the client not in clients_exchange_list, return false. + BuildShareSecretsRsp(share_secrets_resp_builder, schema::ResponseCode_RequestError, + ("client share secret is not in clients_exchange list. && client is illegal"), next_req_time, + iteration); + return false; + } + if (find(clients_share_list.begin(), clients_share_list.end(), fl_id_src) != + clients_share_list.end()) { // the client is already exists, return false. + BuildShareSecretsRsp(share_secrets_resp_builder, schema::ResponseCode_SUCCEED, + ("client sharesecret already exists."), next_req_time, iteration); + return false; + } + + // update new item to memory server. + const flatbuffers::Vector> *encrypted_shares = + (share_secrets_req->encrypted_shares()); + bool retcode_share = cipher_init_->cipher_meta_storage_.UpdateClientShareToServer( + ps::server::kCtxClientsEncryptedShares, fl_id_src, encrypted_shares); + bool retcode_client = + cipher_init_->cipher_meta_storage_.UpdateClientToServer(ps::server::kCtxShareSecretsClientList, fl_id_src); + bool retcode = retcode_share && retcode_client; + if (retcode) { + BuildShareSecretsRsp(share_secrets_resp_builder, schema::ResponseCode_SUCCEED, "OK", next_req_time, iteration); + MS_LOG(INFO) << "CipherShares::ShareSecrets Success"; + } else { + BuildShareSecretsRsp(share_secrets_resp_builder, schema::ResponseCode_OutOfTime, + "update client of shares and shares failed", next_req_time, iteration); + MS_LOG(ERROR) << "CipherShares::ShareSecrets update client of shares and shares failed "; + } + + clock_t end_time = clock(); + double duration = static_cast((end_time - start_time) * 1.0 / CLOCKS_PER_SEC); + MS_LOG(INFO) << "ShareSecrets get + deal + update data time is : " << duration; + return retcode; +} + +bool CipherShares::GetSecrets(const schema::GetShareSecrets *get_secrets_req, + std::shared_ptr get_secrets_resp_builder, + const std::string &next_req_time) { + MS_LOG(INFO) << "CipherShares::GetSecrets START"; + clock_t start_time = clock(); + // step 0: check whether the parameters are legal. + if (get_secrets_req == nullptr) { + BuildGetSecretsRsp(get_secrets_resp_builder, schema::ResponseCode_SystemError, 0, next_req_time, 0); + MS_LOG(ERROR) << "GetSecrets: get_secrets_req is nullptr."; + return false; + } + + // step 1: get client list and client shares list from memory server. + std::vector clients_share_list; + cipher_init_->cipher_meta_storage_.GetClientListFromServer(ps::server::kCtxShareSecretsClientList, + &clients_share_list); + std::map> encrypted_shares_all; + cipher_init_->cipher_meta_storage_.GetClientSharesFromServer(ps::server::kCtxClientsEncryptedShares, + &encrypted_shares_all); + int iteration = get_secrets_req->iteration(); + size_t share_clients_num = clients_share_list.size(); + size_t cients_has_shares = encrypted_shares_all.size(); + if (share_clients_num != cients_has_shares) { + BuildGetSecretsRsp(get_secrets_resp_builder, schema::ResponseCode_OutOfTime, iteration, next_req_time, 0); + MS_LOG(ERROR) << "cients_has_shares: " << cients_has_shares << "share_clients_num: " << share_clients_num; + } + if (cipher_init_->share_clients_num_need_ > share_clients_num) { // the client num is not enough, return false. + BuildGetSecretsRsp(get_secrets_resp_builder, schema::ResponseCode_SucNotReady, iteration, next_req_time, 0); + MS_LOG(INFO) << "GetSecrets: the client num is not enough: share_clients_num_need_: " + << cipher_init_->share_clients_num_need_ << "share_clients_num: " << share_clients_num; + return false; + } + std::string fl_id = get_secrets_req->fl_id()->str(); + if (find(clients_share_list.begin(), clients_share_list.end(), fl_id) == + clients_share_list.end()) { // the client is not in client list, return false. + BuildGetSecretsRsp(get_secrets_resp_builder, schema::ResponseCode_RequestError, iteration, next_req_time, 0); + MS_LOG(ERROR) << "GetSecrets: client is not in client list."; + } + + // get the result client shares. + std::vector encrypted_shares_add; + for (auto encrypted_shares_iterator = encrypted_shares_all.begin(); + encrypted_shares_iterator != encrypted_shares_all.end(); ++encrypted_shares_iterator) { + std::string fl_id_src_now = encrypted_shares_iterator->first; + std::vector &clientshare_str_now = encrypted_shares_iterator->second; + clientshare_str client_share_str_new; + bool find_flag = false; + for (size_t index_clientshare = 0; index_clientshare < clientshare_str_now.size(); ++index_clientshare) { + std::string fl_id_des = clientshare_str_now[index_clientshare].fl_id; + if (fl_id_des == fl_id) { + client_share_str_new.fl_id = fl_id_src_now; + client_share_str_new.index = clientshare_str_now[index_clientshare].index; + client_share_str_new.share = clientshare_str_now[index_clientshare].share; + find_flag = true; + break; + } + } + if (find_flag) { + encrypted_shares_add.push_back(client_share_str_new); + } + } + + // serialise clientshares + size_t size_shares = encrypted_shares_add.size(); + std::vector> encrypted_shares; + std::vector::iterator ptr_start = encrypted_shares_add.begin(); + std::vector::iterator ptr_end = ptr_start + size_shares; + for (std::vector::iterator ptr = ptr_start; ptr < ptr_end; ++ptr) { + auto one_fl_id = get_secrets_resp_builder->CreateString(ptr->fl_id); + auto two_share = get_secrets_resp_builder->CreateVector(ptr->share.data(), ptr->share.size()); + auto third_index = ptr->index; + auto one_clientshare = schema::CreateClientShare(*get_secrets_resp_builder, one_fl_id, two_share, third_index); + encrypted_shares.push_back(one_clientshare); + } + + BuildGetSecretsRsp(get_secrets_resp_builder, schema::ResponseCode_SUCCEED, iteration, next_req_time, + &encrypted_shares); + + MS_LOG(INFO) << "CipherShares::GetSecrets Success"; + clock_t end_time = clock(); + double duration = static_cast((end_time - start_time) * 1.0 / CLOCKS_PER_SEC); + MS_LOG(INFO) << "Getsecrets Duration Time is : " << duration; + return true; +} + +void CipherShares::BuildGetSecretsRsp( + std::shared_ptr get_secrets_resp_builder, schema::ResponseCode retcode, int iteration, + std::string next_req_time, std::vector> *encrypted_shares) { + int rsp_retcode = retcode; + int rsp_iteration = iteration; + auto rsp_next_req_time = get_secrets_resp_builder->CreateString(next_req_time); + if (encrypted_shares == 0) { + auto get_secrets_rsp = + schema::CreateReturnShareSecrets(*get_secrets_resp_builder, rsp_retcode, rsp_iteration, 0, rsp_next_req_time); + get_secrets_resp_builder->Finish(get_secrets_rsp); + } else { + auto encrypted_shares_rsp = get_secrets_resp_builder->CreateVector(*encrypted_shares); + auto get_secrets_rsp = CreateReturnShareSecrets(*get_secrets_resp_builder, rsp_retcode, rsp_iteration, + encrypted_shares_rsp, rsp_next_req_time); + get_secrets_resp_builder->Finish(get_secrets_rsp); + } + + return; +} + +void CipherShares::BuildShareSecretsRsp(std::shared_ptr share_secrets_resp_builder, + const schema::ResponseCode retcode, const string &reason, + const string &next_req_time, const int iteration) { + auto rsp_reason = share_secrets_resp_builder->CreateString(reason); + auto rsp_next_req_time = share_secrets_resp_builder->CreateString(next_req_time); + auto share_secrets_rsp = + schema::CreateResponseShareSecrets(*share_secrets_resp_builder, retcode, rsp_reason, rsp_next_req_time, iteration); + share_secrets_resp_builder->Finish(share_secrets_rsp); + return; +} + +void CipherShares::ClearShareSecrets() { + ps::server::DistributedMetadataStore::GetInstance().ResetMetadata(ps::server::kCtxShareSecretsClientList); + ps::server::DistributedMetadataStore::GetInstance().ResetMetadata(ps::server::kCtxClientsEncryptedShares); +} + +} // namespace armour +} // namespace mindspore diff --git a/mindspore/ccsrc/armour/cipher/cipher_shares.h b/mindspore/ccsrc/armour/cipher/cipher_shares.h new file mode 100644 index 00000000000..add86c44950 --- /dev/null +++ b/mindspore/ccsrc/armour/cipher/cipher_shares.h @@ -0,0 +1,68 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef MINDSPORE_CIPHER_SHARES_H +#define MINDSPORE_CIPHER_SHARES_H + +#include +#include +#include +#include +#include +#include "armour/secure_protocol/secret_sharing.h" +#include "proto/ps.pb.h" +#include "utils/log_adapter.h" +#include "armour/cipher/cipher_init.h" +#include "armour/cipher/cipher_meta_storage.h" + +namespace mindspore { +namespace armour { + +class CipherShares { + public: + // initialize: get cipher_init_ + CipherShares() { cipher_init_ = &CipherInit::GetInstance(); } + + static CipherShares &GetInstance() { + static CipherShares instance; + return instance; + } + + // handle the client's request of share secrets. + bool ShareSecrets(const int cur_iterator, const schema::RequestShareSecrets *share_secrets_req, + std::shared_ptr share_secrets_resp_builder, const string next_req_time); + // handle the client's request of get secrets. + bool GetSecrets(const schema::GetShareSecrets *get_secrets_req, + std::shared_ptr get_secrets_resp_builder, const std::string &next_req_time); + + // build response code of share secrets. + void BuildShareSecretsRsp(std::shared_ptr share_secrets_resp_builder, + const schema::ResponseCode retcode, const string &reason, const string &next_req_time, + const int iteration); + // build response code of get secrets. + void BuildGetSecretsRsp(std::shared_ptr get_secrets_resp_builder, + const schema::ResponseCode retcode, const int iteration, std::string next_req_time, + std::vector> *encrypted_shares); + // clear the shared memory. + void ClearShareSecrets(); + + private: + CipherInit *cipher_init_; // the parameter of the secure aggregation +}; +} // namespace armour +} // namespace mindspore + +#endif // MINDSPORE_CIPHER_KEYS_H diff --git a/mindspore/ccsrc/armour/cipher/cipher_unmask.cc b/mindspore/ccsrc/armour/cipher/cipher_unmask.cc new file mode 100644 index 00000000000..a126f108172 --- /dev/null +++ b/mindspore/ccsrc/armour/cipher/cipher_unmask.cc @@ -0,0 +1,59 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "armour/cipher/cipher_unmask.h" +#include "ps/server/common.h" +#include "ps/server/local_meta_store.h" +#include "armour/cipher/cipher_meta_storage.h" + +namespace mindspore { +namespace armour { +bool CipherUnmask::UnMask(const std::map &data) { + MS_LOG(INFO) << "CipherMgr::UnMask START"; + clock_t start_time = clock(); + std::vector noise; + + cipher_init_->cipher_meta_storage_.GetClientNoisesFromServer(ps::server::kCtxClientNoises, &noise); + if (noise.size() != cipher_init_->featuremap_) { + MS_LOG(ERROR) << " CipherMgr UnMask ERROR"; + return false; + } + + size_t data_size = ps::server::LocalMetaStore::GetInstance().value(ps::server::kCtxFedAvgTotalDataSize); + int sum_size = 0; + for (auto iter = data.begin(); iter != data.end(); ++iter) { + int size_data = iter->second->size / sizeof(float); + float *in_data = reinterpret_cast(iter->second->addr); + MS_LOG(INFO) << " weight name : " << iter->first; + for (int i = 0; i < size_data; ++i) { + in_data[i] = in_data[i] + noise[i + sum_size] / data_size; + } + sum_size += size_data; + for (size_t i = 0; i < data.size(); ++i) { + MS_LOG(INFO) << " index : " << i << " in_data unmask: " << in_data[i] * data_size; + } + } + + MS_LOG(INFO) << "CipherMgr::UnMask sum_size : " << sum_size; + MS_LOG(INFO) << "CipherMgr::UnMask feature_map : " << cipher_init_->featuremap_; + clock_t end_time = clock(); + double duration = static_cast((end_time - start_time) * 1.0 / CLOCKS_PER_SEC); + MS_LOG(INFO) << "Unmask success time is : " << duration; + return true; +} + +} // namespace armour +} // namespace mindspore diff --git a/mindspore/ccsrc/armour/cipher/cipher_unmask.h b/mindspore/ccsrc/armour/cipher/cipher_unmask.h new file mode 100644 index 00000000000..85ddbb421f9 --- /dev/null +++ b/mindspore/ccsrc/armour/cipher/cipher_unmask.h @@ -0,0 +1,45 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef MINDSPORE_CIPHER_UNMASK_H +#define MINDSPORE_CIPHER_UNMASK_H + +#include +#include +#include +#include "armour/secure_protocol/secret_sharing.h" +#include "proto/ps.pb.h" +#include "utils/log_adapter.h" +#include "armour/cipher/cipher_init.h" +#include "armour/cipher/cipher_meta_storage.h" + +namespace mindspore { +namespace armour { + +class CipherUnmask { + public: + // initialize: get cipher_init_ + CipherUnmask() { cipher_init_ = &CipherInit::GetInstance(); } + // unmask the data by secret mask. + bool UnMask(const std::map &data); + + private: + CipherInit *cipher_init_; // the parameter of the secure aggregation +}; +} // namespace armour +} // namespace mindspore + +#endif // MINDSPORE_CIPHER_KEYS_H diff --git a/mindspore/ccsrc/armour/secure_protocol/encrypt.cc b/mindspore/ccsrc/armour/secure_protocol/encrypt.cc new file mode 100644 index 00000000000..b0f5143739b --- /dev/null +++ b/mindspore/ccsrc/armour/secure_protocol/encrypt.cc @@ -0,0 +1,186 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "armour/secure_protocol/encrypt.h" +namespace mindspore { +namespace armour { + +#define KEY_STEP_MAX 32 +#define KEY_STEP_MIN 16 +#define PAD_SIZE 5 + +AESEncrypt::AESEncrypt(const unsigned char *key, int key_len, unsigned char *ivec, int ivec_len, const AES_MODE mode) { + privKey = key; + privKeyLen = key_len; + iVec = ivec; + iVecLen = ivec_len; + aesMode = mode; +} + +AESEncrypt::~AESEncrypt() {} + +int AESEncrypt::EncryptData(const unsigned char *data, const int len, unsigned char *encrypt_data, int *encrypt_len) { + int ret; + if (privKeyLen != KEY_STEP_MIN && privKeyLen != KEY_STEP_MAX) { + std::cout << "key length must be 16 or 32!" << std::endl; + return -1; + } + assert(iVecLen == INIT_VEC_SIZE); + if (aesMode == AES_CBC || aesMode == AES_CTR) { + ret = evp_aes_encrypt(data, len, privKey, iVec, encrypt_data, encrypt_len); + } else { + std::cout << "Please use CBC mode or CTR mode, the other modes are not supported!\n" << std::endl; + ret = -1; + } + if (ret != 0) { + return -1; + } + return 0; +} + +int AESEncrypt::DecryptData(const unsigned char *encrypt_data, const int encrypt_len, unsigned char *data, int *len) { + int ret = 0; + if (privKeyLen != KEY_STEP_MIN && privKeyLen != KEY_STEP_MAX) { + std::cout << "key length must be 16 or 32!" << std::endl; + return -1; + } + assert(iVecLen == INIT_VEC_SIZE); + if (aesMode == AES_CBC || aesMode == AES_CTR) { + ret = evp_aes_decrypt(encrypt_data, encrypt_len, privKey, iVec, data, len); + } else { + std::cout << "Please use CBC mode or CTR mode, the other modes are not supported!" << std::endl; + } + if (ret != 1) { + return -1; + } + return 0; +} + +int AESEncrypt::evp_aes_encrypt(const unsigned char *data, const int len, const unsigned char *key, unsigned char *ivec, + unsigned char *encrypt_data, int *encrypt_len) { + EVP_CIPHER_CTX *ctx = EVP_CIPHER_CTX_new(); + int out_len; + int ret = 0; + if (aesMode == AES_CBC) { + switch (privKeyLen) { + case 16: + ret = EVP_EncryptInit_ex(ctx, EVP_aes_128_cbc(), NULL, key, ivec); + break; + case 32: + ret = EVP_EncryptInit_ex(ctx, EVP_aes_256_cbc(), NULL, key, ivec); + break; + default: + std::cout << "key length is incorrect!" << std::endl; + ret = -1; + } + if (ret != 1) { + std::cout << "EVP_EncryptInit_ex CBC fail!" << std::endl; + return -1; + } + EVP_CIPHER_CTX_set_key_length(ctx, EVP_MAX_KEY_LENGTH); + EVP_CIPHER_CTX_set_padding(ctx, PAD_SIZE); + } else if (aesMode == AES_CTR) { + switch (privKeyLen) { + case 16: + ret = EVP_EncryptInit_ex(ctx, EVP_aes_128_ctr(), NULL, key, ivec); + break; + case 32: + ret = EVP_EncryptInit_ex(ctx, EVP_aes_256_ctr(), NULL, key, ivec); + break; + default: + std::cout << "key length is incorrect!" << std::endl; + ret = -1; + } + if (ret != 1) { + std::cout << "EVP_EncryptInit_ex CTR fail!" << std::endl; + return -1; + } + } else { + std::cout << "Unsupported AES mode" << std::endl; + return -1; + } + ret = EVP_EncryptUpdate(ctx, encrypt_data, &out_len, data, len); + if (ret != 1) { + std::cout << "EVP_EncryptUpdate fail!" << std::endl; + return -1; + } + *encrypt_len = out_len; + ret = EVP_EncryptFinal_ex(ctx, encrypt_data + *encrypt_len, &out_len); + if (ret != 1) { + std::cout << "EVP_EncryptFinal_ex fail!" << std::endl; + return -1; + } + *encrypt_len += out_len; + EVP_CIPHER_CTX_free(ctx); + return 0; +} + +int AESEncrypt::evp_aes_decrypt(const unsigned char *encrypt_data, const int len, const unsigned char *key, + unsigned char *ivec, unsigned char *decrypt_data, int *decrypt_len) { + EVP_CIPHER_CTX *ctx = EVP_CIPHER_CTX_new(); + int out_len; + int ret = 0; + if (aesMode == AES_CBC) { + switch (privKeyLen) { + case 16: + ret = EVP_DecryptInit_ex(ctx, EVP_aes_128_cbc(), NULL, key, ivec); + break; + case 32: + ret = EVP_DecryptInit_ex(ctx, EVP_aes_256_cbc(), NULL, key, ivec); + break; + default: + std::cout << "key length is incorrect!" << std::endl; + ret = -1; + } + if (ret != 1) { + return -1; + } + EVP_CIPHER_CTX_set_key_length(ctx, EVP_MAX_KEY_LENGTH); + } else if (aesMode == AES_CTR) { + switch (privKeyLen) { + case 16: + ret = EVP_DecryptInit_ex(ctx, EVP_aes_128_ctr(), NULL, key, ivec); + break; + case 32: + ret = EVP_DecryptInit_ex(ctx, EVP_aes_256_ctr(), NULL, key, ivec); + break; + default: + std::cout << "key length is incorrect!" << std::endl; + ret = -1; + } + } else { + ret = -1; + } + + if (ret != 1) { + return -1; + } + ret = EVP_DecryptUpdate(ctx, decrypt_data, &out_len, encrypt_data, len); + if (ret != 1) { + return -1; + } + *decrypt_len = out_len; + ret = EVP_DecryptFinal_ex(ctx, decrypt_data + *decrypt_len, &out_len); + if (ret != 1) { + return -1; + } + *decrypt_len += out_len; + EVP_CIPHER_CTX_free(ctx); + + return 0; +} +} // namespace armour +} // namespace mindspore diff --git a/mindspore/ccsrc/armour/secure_protocol/encrypt.h b/mindspore/ccsrc/armour/secure_protocol/encrypt.h new file mode 100644 index 00000000000..d0a996ecd9c --- /dev/null +++ b/mindspore/ccsrc/armour/secure_protocol/encrypt.h @@ -0,0 +1,61 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef MINDSPORE_ARMOUR_ENCRYPT_H +#define MINDSPORE_ARMOUR_ENCRYPT_H + +#include +#include +#include + +#define INIT_VEC_SIZE 16 + +namespace mindspore { +namespace armour { +class Encrypt {}; +enum AES_MODE { + AES_CBC = 0, + AES_CTR = 1, +}; +class SymmetricEncrypt : Encrypt {}; + +class AESEncrypt : SymmetricEncrypt { + // use openssl EVP_aes_256_cbc/EVP_aes_128_ctr + // hash input key to fixed-length (128/256 bits) using md5/SHA-256 + public: + AESEncrypt(const unsigned char *key, int key_len, unsigned char *ivec, int ivec_len, AES_MODE mode); + ~AESEncrypt(); + int EncryptData(const unsigned char *data, const int len, unsigned char *encrypt_data, int *encrypt_len); + int DecryptData(const unsigned char *encrypt_data, const int encrypt_len, unsigned char *data, int *len); + + private: + const unsigned char *privKey; + int privKeyLen; + unsigned char *iVec; + int iVecLen; + AES_MODE aesMode; + // int evp_aes_cbc_encrypt(const unsigned char* data, const int len, const unsigned char* key, unsigned char* ivec, + // unsigned char* encrypt_data, int& encrypt_len); int evp_aes_cbc_decrypt(const unsigned char* encrypt_data, const + // int len, const unsigned char* key, unsigned char* ivec, unsigned char* decrypt_data, int& decrypt_len); + int evp_aes_encrypt(const unsigned char *data, const int len, const unsigned char *key, unsigned char *ivec, + unsigned char *encrypt_data, int *encrypt_len); + int evp_aes_decrypt(const unsigned char *encrypt_data, const int len, const unsigned char *key, unsigned char *ivec, + unsigned char *decrypt_data, int *decrypt_len); +}; + +} // namespace armour +} // namespace mindspore +#endif // MINDSPORE_ARMOUR_ENCRYPT_H diff --git a/mindspore/ccsrc/armour/secure_protocol/key_agreement.cc b/mindspore/ccsrc/armour/secure_protocol/key_agreement.cc new file mode 100644 index 00000000000..0df887fe06e --- /dev/null +++ b/mindspore/ccsrc/armour/secure_protocol/key_agreement.cc @@ -0,0 +1,146 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "armour/secure_protocol/key_agreement.h" +#include + +namespace mindspore { +namespace armour { +PublicKey::PublicKey(EVP_PKEY *evpKey) { evpPubKey = evpKey; } + +PublicKey::~PublicKey() { EVP_PKEY_free(evpPubKey); } + +PrivateKey::PrivateKey(EVP_PKEY *evpKey) { evpPrivKey = evpKey; } + +PrivateKey::~PrivateKey() { EVP_PKEY_free(evpPrivKey); } + +int PrivateKey::GetPrivateBytes(size_t *len, unsigned char *privKeyBytes) { + if (!EVP_PKEY_get_raw_private_key(evpPrivKey, privKeyBytes, len)) { + return -1; + } + return 0; +} + +int PrivateKey::GetPublicBytes(size_t *len, unsigned char *pubKeyBytes) { + if (!EVP_PKEY_get_raw_public_key(evpPrivKey, pubKeyBytes, len)) { + return -1; + } + return 0; +} + +int PrivateKey::Exchange(PublicKey *peerPublicKey, int key_len, const unsigned char *salt, int salt_len, + unsigned char *exchangeKey) { + EVP_PKEY_CTX *ctx; + size_t len = 0; + ctx = EVP_PKEY_CTX_new(evpPrivKey, NULL); + if (!ctx) { + std::cout << "EVP_PKEY_CTX_new failed!" << std::endl; + return -1; + } + if (EVP_PKEY_derive_init(ctx) <= 0) { + std::cout << "EVP_PKEY_derive_init failed!" << std::endl; + return -1; + } + if (EVP_PKEY_derive_set_peer(ctx, peerPublicKey->evpPubKey) <= 0) { + std::cout << "EVP_PKEY_derive_set_peer failed!" << std::endl; + return -1; + } + unsigned char *secret; + if (EVP_PKEY_derive(ctx, NULL, &len) <= 0) { + std::cout << "get derive key size failed!" << std::endl; + return -1; + } + + secret = (unsigned char *)OPENSSL_malloc(len); + if (!secret) { + std::cout << "malloc secret memory failed!" << std::endl; + return -1; + } + + if (EVP_PKEY_derive(ctx, secret, &len) <= 0) { + std::cout << "derive key failed!" << std::endl; + return -1; + } + + if (!PKCS5_PBKDF2_HMAC((char *)secret, len, salt, salt_len, ITERATION, EVP_sha256(), key_len, exchangeKey)) { + return -1; + } + OPENSSL_free(secret); + EVP_PKEY_CTX_free(ctx); + return 0; +} + +// using x25519 curve +PrivateKey *KeyAgreement::GeneratePrivKey() { + EVP_PKEY *evpKey = NULL; + EVP_PKEY_CTX *pctx = EVP_PKEY_CTX_new_id(EVP_PKEY_X25519, NULL); + if (!pctx) { + return NULL; + } + if (EVP_PKEY_keygen_init(pctx) <= 0) { + return NULL; + } + if (EVP_PKEY_keygen(pctx, &evpKey) <= 0) { + return NULL; + } + EVP_PKEY_CTX_free(pctx); + // PEM_write_PrivateKey(stdout, evpKey, NULL, NULL, 0, NULL, NULL); + PrivateKey *privKey = new PrivateKey(evpKey); + return privKey; +} + +PublicKey *KeyAgreement::GeneratePubKey(PrivateKey *privKey) { + unsigned char *pubKeyBytes; + size_t len = 0; + if (!EVP_PKEY_get_raw_public_key(privKey->evpPrivKey, NULL, &len)) { + return NULL; + } + pubKeyBytes = (unsigned char *)OPENSSL_malloc(len); + if (!EVP_PKEY_get_raw_public_key(privKey->evpPrivKey, pubKeyBytes, &len)) { + return NULL; + } + EVP_PKEY *evp_pubKey = EVP_PKEY_new_raw_public_key(EVP_PKEY_X25519, NULL, (unsigned char *)pubKeyBytes, len); + OPENSSL_free(pubKeyBytes); + PublicKey *pubKey = new PublicKey(evp_pubKey); + return pubKey; +} + +PrivateKey *KeyAgreement::FromPrivateBytes(unsigned char *data, int len) { + EVP_PKEY *evp_Key = EVP_PKEY_new_raw_private_key(EVP_PKEY_X25519, NULL, data, len); + if (evp_Key == NULL) { + return NULL; + } + PrivateKey *privKey = new PrivateKey(evp_Key); + return privKey; +} + +PublicKey *KeyAgreement::FromPublicBytes(unsigned char *data, int len) { + EVP_PKEY *evp_pubKey = EVP_PKEY_new_raw_public_key(EVP_PKEY_X25519, NULL, data, len); + if (evp_pubKey == NULL) { + std::cout << "create evp_pubKey from raw bytes fail" << std::endl; + return NULL; + } + PublicKey *pubKey = new PublicKey(evp_pubKey); + return pubKey; +} + +int KeyAgreement::ComputeSharedKey(PrivateKey *privKey, PublicKey *peerPublicKey, int key_len, + const unsigned char *salt, int salt_len, unsigned char *exchangeKey) { + return privKey->Exchange(peerPublicKey, key_len, salt, salt_len, exchangeKey); +} + +} // namespace armour +} // namespace mindspore diff --git a/mindspore/ccsrc/armour/secure_protocol/key_agreement.h b/mindspore/ccsrc/armour/secure_protocol/key_agreement.h new file mode 100644 index 00000000000..895604ed812 --- /dev/null +++ b/mindspore/ccsrc/armour/secure_protocol/key_agreement.h @@ -0,0 +1,60 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef MINDSPORE_KEY_AGREEMENT_H +#define MINDSPORE_KEY_AGREEMENT_H + +#include +#include +#include +#include + +#define KEK_KEY_LEN 32 +#define ITERATION 10000 + +namespace mindspore { +namespace armour { +class PublicKey { + public: + explicit PublicKey(EVP_PKEY *evpKey); + ~PublicKey(); + EVP_PKEY *evpPubKey; +}; + +class PrivateKey { + public: + explicit PrivateKey(EVP_PKEY *evpKey); + ~PrivateKey(); + int Exchange(PublicKey *peerPublicKey, int key_len, const unsigned char *salt, int salt_len, + unsigned char *exchangeKey); + int GetPrivateBytes(size_t *len, unsigned char *priKeyBytes); + int GetPublicBytes(size_t *len, unsigned char *pubKeyBytes); + EVP_PKEY *evpPrivKey; +}; + +class KeyAgreement { + public: + static PrivateKey *GeneratePrivKey(); + static PublicKey *GeneratePubKey(PrivateKey *privKey); + static PrivateKey *FromPrivateBytes(unsigned char *data, int len); + static PublicKey *FromPublicBytes(unsigned char *data, int len); + static int ComputeSharedKey(PrivateKey *privKey, PublicKey *peerPublicKey, int key_len, const unsigned char *salt, + int salt_len, unsigned char *exchangeKey); +}; + +} // namespace armour +} // namespace mindspore +#endif // MINDSPORE_KEY_AGREEMENT_H diff --git a/mindspore/ccsrc/armour/secure_protocol/random.cc b/mindspore/ccsrc/armour/secure_protocol/random.cc new file mode 100644 index 00000000000..b17bdc2fc19 --- /dev/null +++ b/mindspore/ccsrc/armour/secure_protocol/random.cc @@ -0,0 +1,74 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#include "armour/secure_protocol/random.h" +#include +namespace mindspore { +namespace armour { +Random::Random(size_t init_seed) { generator.seed(init_seed); } + +Random::~Random() {} + +int Random::GetRandomBytes(unsigned char *secret, int num_bytes) { + int retval = RAND_priv_bytes(secret, RANDOM_LEN); + return retval; +} + +void Random::RandUniform(float *array, int size) { + std::uniform_real_distribution rand(0, 1); + for (int i = 0; i < size; i++) { + *(reinterpret_cast(array) + i) = rand(generator); + } +} + +void Random::RandNorminal(float *array, int size) { + std::normal_distribution randn(0, 1); + for (int i = 0; i < size; i++) { + *(reinterpret_cast(array) + i) = randn(generator); + } +} + +int Random::RandomAESCTR(std::vector *noise, int noise_len, const unsigned char *seed, int seed_len) { + if (seed_len != 16 && seed_len != 32) { + std::cout << "seed length must be 16 or 32!" << std::endl; + return -1; + } + int size = noise_len * sizeof(int); + unsigned char data[size]; + unsigned char encrypt_data[size]; + for (int i = 0; i < size; i++) { + data[i] = 0; + encrypt_data[i] = 0; + } + unsigned char ivec[INIT_VEC_SIZE]; + for (size_t i = 0; i < INIT_VEC_SIZE; i++) { + ivec[i] = 0; + } + int encrypt_len; + AESEncrypt encrypt(seed, seed_len, ivec, INIT_VEC_SIZE, AES_CTR); + if (encrypt.EncryptData(data, size, encrypt_data, &encrypt_len) != 0) { + std::cout << "call encryptData fail!" << std::endl; + return -1; + } + + for (int i = 0; i < noise_len; i++) { + auto value = *(reinterpret_cast(encrypt_data) + i); + noise->emplace_back(static_cast(value) / INT32_MAX); + } + return 0; +} +} // namespace armour +} // namespace mindspore diff --git a/mindspore/ccsrc/armour/secure_protocol/random.h b/mindspore/ccsrc/armour/secure_protocol/random.h new file mode 100644 index 00000000000..ae6cfb3d567 --- /dev/null +++ b/mindspore/ccsrc/armour/secure_protocol/random.h @@ -0,0 +1,47 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +#ifndef MINDSPORE_ARMOUR_RANDOM_H +#define MINDSPORE_ARMOUR_RANDOM_H +#include +#include +#include +#include +#include "armour/secure_protocol/encrypt.h" +namespace mindspore { +namespace armour { + +#define RANDOM_LEN 8 + +class Random { + public: + explicit Random(size_t init_seed); + ~Random(); + // use openssl RAND_priv_bytes + static int GetRandomBytes(unsigned char *secret, int num_bytes); + + // std::uniform_real_distribution rand(0,1) + void RandUniform(float *array, int size); + // std::normal_distribution randn(0,1); + void RandNorminal(float *array, int size); + static int RandomAESCTR(std::vector *noise, int noise_len, const unsigned char *seed, int seed_len); + + private: + std::default_random_engine generator; +}; +} // namespace armour +} // namespace mindspore +#endif // MINDSPORE_ARMOUR_RANDOM_H diff --git a/mindspore/ccsrc/armour/secure_protocol/secret_sharing.cc b/mindspore/ccsrc/armour/secure_protocol/secret_sharing.cc new file mode 100644 index 00000000000..22b8cc5f260 --- /dev/null +++ b/mindspore/ccsrc/armour/secure_protocol/secret_sharing.cc @@ -0,0 +1,220 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#include "armour/secure_protocol/secret_sharing.h" +#include +namespace mindspore { +namespace armour { +void secure_zero(unsigned char *s, size_t n) { + volatile unsigned char *p = s; + if (p) + while (n--) *p++ = '\0'; +} + +int GetRandInteger(mpz_t x, mpz_t prim) { + size_t bytes_len = (mpz_sizeinbase(prim, 2) + 8 - 1) / 8; + unsigned char buf[bytes_len]; + while (true) { + if (!RAND_bytes(buf, bytes_len)) { + std::cout << "Get Rand Integer failed!" << std::endl; + continue; + } + mpz_import(x, bytes_len, 1, 1, 0, 0, buf); + secure_zero(buf, sizeof(buf)); + if (mpz_cmp_ui(x, 0) > 0 && mpz_cmp(x, prim) < 0) { + return 0; + } + } +} + +int GetRandomPrime(mpz_t prim) { + mpz_t rand; + mpz_init(rand); + const int max_prime_len = SECRET_MAX_LEN + 1; + unsigned char buf[max_prime_len]; + if (!RAND_bytes(buf, max_prime_len)) { + std::cout << "Get Rand Integer failed!" << std::endl; + return -1; + } + mpz_import(rand, max_prime_len, 1, 1, 0, 0, buf); + mpz_nextprime(prim, rand); + mpz_clear(rand); + secure_zero(buf, sizeof(buf)); + return 0; +} + +void PrintBigInteger(mpz_t x) { + char *tmp = mpz_get_str(NULL, 16, x); + std::string Str = tmp; + std::cout << "*************************" << Str << std::endl; + void (*freefunc)(void *, size_t); + mp_get_memory_functions(NULL, NULL, &freefunc); + freefunc(tmp, strlen(tmp) + 1); +} + +void PrintBigInteger(mpz_t x, int hex) { + char *tmp = mpz_get_str(NULL, hex, x); + std::string Str = tmp; + + std::cout << Str << std::endl; + void (*freefunc)(void *, size_t); + mp_get_memory_functions(NULL, NULL, &freefunc); + freefunc(tmp, strlen(tmp) + 1); +} + +Share::~Share() { + if (this->data != nullptr) free(this->data); +} + +SecretSharing::SecretSharing(mpz_t prim) { + mpz_init(this->prim_); + mpz_set(this->prim_, prim); +} + +SecretSharing::~SecretSharing() { mpz_clear(this->prim_); } + +void SecretSharing::GetPolyVal(int k, mpz_t y, const mpz_t x, const mpz_t coeff[]) { + int i; + mpz_set_ui(y, 0); + for (i = k - 1; i >= 0; i--) { + field_mult(y, y, x); + field_add(y, y, coeff[i]); + } +} + +void SecretSharing::field_invert(mpz_t z, const mpz_t x) { mpz_invert(z, x, this->prim_); } + +void SecretSharing::field_add(mpz_t z, const mpz_t x, const mpz_t y) { + mpz_add(z, x, y); + mpz_mod(z, z, this->prim_); +} + +void SecretSharing::field_mult(mpz_t z, const mpz_t x, const mpz_t y) { + mpz_mul(z, x, y); + mpz_mod(z, z, this->prim_); +} + +int SecretSharing::CalculateShares(const mpz_t coeff[], int k, int n, const std::vector &shares) { + mpz_t x, y; + mpz_init(x); + mpz_init(y); + for (int i = 0; i < n; i++) { + mpz_set_ui(x, i + 1); + GetPolyVal(k, y, x, coeff); + shares[i]->index = i + 1; + size_t share_len = (mpz_sizeinbase(y, 2) + 8 - 1) / 8; + shares[i]->data = (unsigned char *)malloc(share_len + 1); + mpz_export(shares[i]->data, &(shares[i]->len), 1, 1, 0, 0, y); + if (shares[i]->len != share_len) { + std::cout << "share_len is not equal" << std::endl; + return -1; + } + std::cout << "share_" << i + 1 << ": "; + PrintBigInteger(y); + } + mpz_clear(x); + mpz_clear(y); + return 0; +} + +int SecretSharing::Split(int n, const int k, const char *secret, const size_t length, + const std::vector &shares) { + if (k <= 1 || k > n) { + std::cout << "invalid parameters" << std::endl; + return -1; + } + if (static_cast(shares.size()) != n) { + std::cout << "the size of shares must be equal to nq" << std::endl; + return -1; + } + this->degree_ = length * 8; + const int kCoeffLen = k; + mpz_t coeff[kCoeffLen]; + int ret = 0; + int i = 0; + mpz_init(coeff[i]); + mpz_import(coeff[i], length, 1, 1, 0, 0, secret); + i++; + for (; i < k && ret == 0; i++) { + mpz_init(coeff[i]); + ret = GetRandInteger(coeff[i], this->prim_); + std::cout << "coeff_" << i << ":"; + PrintBigInteger(coeff[i]); + } + if (ret == 0) ret = CalculateShares(coeff, k, n, shares); + for (i = 0; i < k; i++) mpz_clear(coeff[i]); + return ret; +} + +void SecretSharing::GetShare(mpz_t x, mpz_t share, Share *s_share) { + mpz_set_ui(x, s_share->index); + mpz_import(share, s_share->len, 1, 1, 0, 0, s_share->data); +} + +int SecretSharing::Combine(int k, const std::vector &shares, char *secret, size_t *length) { + int ret = 0; + mpz_t y[k], x[k], denses[k], nums[k]; + int i, j, m; + + for (i = 0; i < k; i++) { + mpz_init(x[i]); + mpz_init(y[i]); + mpz_init(denses[i]); + mpz_init(nums[i]); + GetShare(x[i], y[i], shares[i]); + std::cout << "combine -- share_" << mpz_get_str(NULL, 10, x[i]) << ": "; + PrintBigInteger(y[i]); + printf("index is : %u\n", shares[i]->index); + printf("len is %zu.\n", shares[i]->len); + } + + mpz_t sum; + mpz_init(sum); + mpz_set_ui(sum, 0); + for (j = 0; j < k; j++) { + mpz_set_ui(denses[j], 1); + mpz_set_ui(nums[j], 1); + mpz_t tmp; + mpz_init(tmp); + for (m = 0; m < k; m++) { + if (m != j) { + field_mult(nums[j], nums[j], x[m]); + mpz_mul_si(tmp, x[j], -1); + field_add(tmp, x[m], tmp); + field_mult(denses[j], denses[j], tmp); + } + } + field_invert(tmp, denses[j]); + + field_mult(tmp, tmp, nums[j]); + field_mult(tmp, tmp, y[j]); + field_add(sum, sum, tmp); + mpz_clear(tmp); + } + + mpz_export(secret, length, 1, 1, 0, 0, sum); + PrintBigInteger(sum); + mpz_clear(sum); + for (i = 0; i < k; i++) { + mpz_clear(x[i]); + mpz_clear(y[i]); + mpz_clear(nums[i]); + mpz_clear(denses[i]); + } + return ret; +} + +} // namespace armour +} // namespace mindspore diff --git a/mindspore/ccsrc/armour/secure_protocol/secret_sharing.h b/mindspore/ccsrc/armour/secure_protocol/secret_sharing.h new file mode 100644 index 00000000000..a19221e8853 --- /dev/null +++ b/mindspore/ccsrc/armour/secure_protocol/secret_sharing.h @@ -0,0 +1,72 @@ +/** + * Copyright 2021 Huawei Technologies Co., Ltd + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ +#ifndef MINDSPORE_SECRET_SHARING_H +#define MINDSPORE_SECRET_SHARING_H +#include +#include +#include +#include +#include +#include +#include +#include "openssl/rand.h" + +namespace mindspore { +namespace armour { +#define SECRET_MAX_LEN 32 +#define PRIME_MAX_LEN 33 + +struct Share { + unsigned int index; + unsigned char *data; + size_t len; + ~Share(); +}; + +void secure_zero(void *s, size_t); +int GetRandInteger(mpz_t x, mpz_t prim); +int GetRandomPrime(mpz_t prim); +void PrintBigInteger(mpz_t x); +void PrintBigInteger(mpz_t x, int hex); + +class SecretSharing { + public: + explicit SecretSharing(mpz_t prim); + ~SecretSharing(); + // split the input secret into multiple shares + int Split(int n, const int k, const char *secret, size_t length, const std::vector &shares); + // reconstruct the secret from multiple shares + int Combine(int k, const std::vector &shares, char *secret, size_t *length); + + private: + mpz_t prim_; + size_t degree_; + // calculate shares from a polynomial + int CalculateShares(const mpz_t coeff[], int k, int n, const std::vector &shares); + // inversion in finite field + void field_invert(mpz_t z, const mpz_t x); + // addition in finite field + void field_add(mpz_t z, const mpz_t x, const mpz_t y); + // multiplication in finite field + void field_mult(mpz_t z, const mpz_t x, const mpz_t y); + // evaluate polynomial at x + void GetPolyVal(int k, mpz_t y, const mpz_t x, const mpz_t coeff[]); + // convert secret sharing from Share type to mpz_t type + void GetShare(mpz_t x, mpz_t share, Share *s_share); +}; +} // namespace armour +} // namespace mindspore +#endif // MINDSPORE_SECRET_SHARING_H