diff --git a/mindspore_serving/ccsrc/common/tensor.cc b/mindspore_serving/ccsrc/common/tensor.cc index 51aa6d1..cea636d 100644 --- a/mindspore_serving/ccsrc/common/tensor.cc +++ b/mindspore_serving/ccsrc/common/tensor.cc @@ -14,10 +14,10 @@ * limitations under the License. */ #include "common/tensor.h" +#include #include #include #include "common/log.h" -#include "securec.h" namespace mindspore::serving { diff --git a/mindspore_serving/ccsrc/worker/worker.cc b/mindspore_serving/ccsrc/worker/worker.cc index be810e1..050050e 100644 --- a/mindspore_serving/ccsrc/worker/worker.cc +++ b/mindspore_serving/ccsrc/worker/worker.cc @@ -152,7 +152,7 @@ Status Worker::Run(const RequestSpec &request_spec, const std::vectorHasNext()) { serving::Instance instance; - result->GetNext(instance); + result->GetNext(&instance); outputs->push_back(instance); } return SUCCESS; @@ -395,8 +395,7 @@ void Worker::GetVersions(const LoadServableSpec &servable_spec, std::vector= future_list_.size()) { MSI_LOG_ERROR << "GetNext failed, index greater than instance count " << future_list_.size(); return FAILED; @@ -497,14 +497,14 @@ Status AsyncResult::GetNext(Instance &instance_result) { next_index_++; auto &future = future_list_[index]; if (!future.valid()) { - instance_result.error_msg = result_[index].error_msg; + instance_result->error_msg = result_[index].error_msg; return FAILED; } const int kWaitMaxHundredMs = 100; int i; for (i = 0; i < kWaitMaxHundredMs; i++) { // if (Worker::GetInstance().HasCleared()) { - instance_result.error_msg = Status(FAILED, "Servable stopped"); + instance_result->error_msg = Status(FAILED, "Servable stopped"); return FAILED; } if (future.wait_for(std::chrono::milliseconds(100)) == std::future_status::ready) { @@ -513,12 +513,12 @@ Status AsyncResult::GetNext(Instance &instance_result) { } if (i >= kWaitMaxHundredMs) { MSI_LOG_ERROR << "GetNext failed, wait time out, index " << index << ", total count " << future_list_.size(); - instance_result.error_msg = Status(FAILED, "Time out"); + instance_result->error_msg = Status(FAILED, "Time out"); return FAILED; } future.get(); - instance_result = result_[index]; + *instance_result = result_[index]; return SUCCESS; } diff --git a/mindspore_serving/ccsrc/worker/worker.h b/mindspore_serving/ccsrc/worker/worker.h index b86e028..d30ad4f 100644 --- a/mindspore_serving/ccsrc/worker/worker.h +++ b/mindspore_serving/ccsrc/worker/worker.h @@ -38,7 +38,7 @@ class AsyncResult { explicit AsyncResult(size_t size); bool HasNext(); - Status GetNext(Instance &instance_result); + Status GetNext(Instance *instance_result); private: std::vector> future_list_; diff --git a/mindspore_serving/client/cpp/client.cc b/mindspore_serving/client/cpp/client.cc index 168edad..a7f823f 100644 --- a/mindspore_serving/client/cpp/client.cc +++ b/mindspore_serving/client/cpp/client.cc @@ -420,11 +420,14 @@ class ClientImpl { auto channel = grpc::CreateChannel(target_str, grpc::InsecureChannelCredentials()); stub_ = proto::MSService::NewStub(channel); } - Status Predict(const proto::PredictRequest &request, proto::PredictReply &reply) { + Status Predict(const proto::PredictRequest &request, proto::PredictReply *reply) { + if (reply == nullptr) { + return Status(SYSTEM_ERROR, "ClientImpl::Predict input reply cannot be nullptr"); + } grpc::ClientContext context; // The actual RPC. - grpc::Status status = stub_->Predict(&context, request, &reply); + grpc::Status status = stub_->Predict(&context, request, reply); if (status.ok()) { return SUCCESS; } else { @@ -457,10 +460,7 @@ Status Client::SendRequest(const InstancesRequest &request, InstancesReply *repl servable_spec->set_method_name(method_name_); servable_spec->set_version_number(version_number_); - Status result = impl_->Predict(*proto_request, *proto_reply); - // std::string str; - // google::protobuf::TextFormat::PrintToString(*proto_reply, &str); - // std::cout << str << std::endl; + Status result = impl_->Predict(*proto_request, proto_reply); return result; }