From c460b0541b02cfc23323c719f617188266876a91 Mon Sep 17 00:00:00 2001 From: Sourabh Betigeri Date: Thu, 23 Jan 2025 16:22:13 -0800 Subject: [PATCH] SWDEV-502219 - Adds validity checks for negative parameters passed Change-Id: Ib8a531533306a27143d74b81c074de81051eb896 --- hipamd/src/hip_graph.cpp | 16 ++++++++++++++++ hipamd/src/hip_stream_ops.cpp | 10 +++++++++- 2 files changed, 25 insertions(+), 1 deletion(-) diff --git a/hipamd/src/hip_graph.cpp b/hipamd/src/hip_graph.cpp index b517ddbb7f..f6bd51450b 100644 --- a/hipamd/src/hip_graph.cpp +++ b/hipamd/src/hip_graph.cpp @@ -3563,6 +3563,12 @@ hipError_t hipGraphAddBatchMemOpNode(hipGraphNode_t* phGraphNode, hipGraph_t hGr (numDependencies > 0 && dependencies == nullptr) || nodeParams == nullptr) { HIP_RETURN(hipErrorInvalidValue); } + // Check nodeParams fields + if (nodeParams->count <= 0 || nodeParams->count > 256 || nodeParams->paramArray == nullptr || + nodeParams->flags != 0 || nodeParams->ctx == nullptr) { + return hipErrorInvalidValue; + } + hip::GraphNode* node = new hip::hipGraphBatchMemOpNode(nodeParams); hipError_t status = ihipGraphAddNode(node, reinterpret_cast(hGraph), @@ -3589,6 +3595,11 @@ hipError_t hipGraphBatchMemOpNodeSetParams(hipGraphNode_t hNode, if (!hip::GraphNode::isNodeValid(n) || nodeParams == nullptr) { HIP_RETURN(hipErrorInvalidValue); } + // Check nodeParams fields + if (nodeParams->count <= 0 || nodeParams->count > 256 || nodeParams->paramArray == nullptr || + nodeParams->flags != 0 || nodeParams->ctx == nullptr) { + return hipErrorInvalidValue; + } HIP_RETURN(reinterpret_cast(n)->SetParams(nodeParams)); } @@ -3602,6 +3613,11 @@ hipError_t hipGraphExecBatchMemOpNodeSetParams(hipGraphExec_t hGraphExec, !hip::GraphNode::isNodeValid(n) || nodeParams == nullptr) { HIP_RETURN(hipErrorInvalidValue); } + // Check nodeParams fields + if (nodeParams->count <= 0 || nodeParams->count > 256 || nodeParams->paramArray == nullptr || + nodeParams->flags != 0 || nodeParams->ctx == nullptr) { + return hipErrorInvalidValue; + } hip::GraphNode* clonedNode = reinterpret_cast(graphExec)->GetClonedNode(n); if (clonedNode == nullptr) { HIP_RETURN(hipErrorInvalidValue); diff --git a/hipamd/src/hip_stream_ops.cpp b/hipamd/src/hip_stream_ops.cpp index 3287d47d49..952e2f2b19 100644 --- a/hipamd/src/hip_stream_ops.cpp +++ b/hipamd/src/hip_stream_ops.cpp @@ -25,7 +25,7 @@ namespace hip { hipError_t ihipBatchMemOperation(hipStream_t stream, cl_command_type cmdType, unsigned int count, hipStreamBatchMemOpParams* paramArray, unsigned int flags) { - if (paramArray == nullptr || flags != 0 || count > 256) { + if (paramArray == nullptr || flags != 0 || count == 0 || count > 256) { return hipErrorInvalidValue; } @@ -33,6 +33,14 @@ hipError_t ihipBatchMemOperation(hipStream_t stream, cl_command_type cmdType, un return hipErrorContextIsDestroyed; } + // Validate operations in paramArray + for (unsigned int i = 0; i < count; i++) { + // These operations are currently not supported + if (paramArray[i].operation == hipStreamMemOpBarrier || hipStreamMemOpFlushRemoteWrites) { + return hipErrorInvalidValue; + } + } + hip::Stream* hip_stream = hip::getStream(stream); amd::Command::EventWaitList waitList;