* Copyright (c) 2026 Huawei Technologies Co., Ltd.
* This program is free software, you can redistribute it and/or modify it under the terms and conditions of
* CANN Open Software License Agreement Version 2.0 (the "License").
* Please refer to the License for details. You may not use this file except in compliance with the License.
* THIS SOFTWARE IS PROVIDED ON AN "AS IS" BASIS, WITHOUT WARRANTIES OF ANY KIND, EITHER EXPRESS OR IMPLIED,
* INCLUDING BUT NOT LIMITED TO NON-INFRINGEMENT, MERCHANTABILITY, OR FITNESS FOR A PARTICULAR PURPOSE.
* See LICENSE in the root of the software repository for the full text of the License.
*/
#include "api_handle_guard.h"
#include "api.hpp"
#include "api_global_err.h"
#include "api_kernel_func.hpp"
#include "args/args_inner.h"
#include "errcode_manage.hpp"
#include "error_message_manage.hpp"
#include "spec/base_info.hpp"
namespace cce {
namespace runtime {
namespace {
rtError_t ReportApiHandleValidationError(const rtError_t errCode, const char_t* callerFuncName)
{
const std::string errorStr = RT_GET_ERRDESC(errCode);
RT_LOG(RT_LOG_ERROR, "%s", errorStr.c_str());
ErrorMessageUtils::FuncErrorReason(errCode, callerFuncName);
RT_LOG_FLUSH();
return GetRtExtErrCodeAndSetGlobalErr(errCode);
}
static bool HasExpectedMagic(const void* handle, const uint64_t expectedMagic)
{
const auto* const innerObject = static_cast<const rtInnerObject*>(handle);
return innerObject->magic.load(std::memory_order_acquire) == expectedMagic;
}
template <typename T>
rtError_t ProbeValidatedObject(const void* handle, T*& outRealObj)
{
if (handle == nullptr) {
outRealObj = nullptr;
return RT_ERROR_NONE;
}
const auto* const innerObject = static_cast<const rtInnerObject*>(handle);
if (innerObject->magic.load(std::memory_order_acquire) != RtMagicTraits<T>::value) {
return RT_ERROR_INVALID_HANDLE;
}
outRealObj = static_cast<T*>(innerObject->object);
return RT_ERROR_NONE;
}
template <typename ApiType>
rtError_t ConvertFuncToKernelImpl(
ApiType* const apiInstance, const void* func, Kernel*& kernel, const char_t* const callerFuncName)
{
Kernel* symbolKernel = nullptr;
const rtError_t ret = apiInstance->GetFunctionBySymbol(func, &symbolKernel);
if (ret == RT_ERROR_NONE) {
RT_LOG(RT_LOG_INFO, "find function handle by symbol");
kernel = symbolKernel;
return RT_ERROR_NONE;
} else if (ret == RT_ERROR_INVALID_DEVICE_FUNCTION) {
RT_LOG(RT_LOG_INFO, "cannot find function handle by symbol, treat func as handle directly");
return ValidateKernelHandleForApi(func, kernel, callerFuncName);
} else {
ErrorMessageUtils::FuncErrorReason(ret, callerFuncName);
RT_LOG_FLUSH();
return GetRtExtErrCodeAndSetGlobalErr(ret);
}
}
}
rtError_t ValidateModelHandleForApi(rtModel_t handle, Model*& outRealObj, const char_t* callerFuncName)
{
const rtError_t ret = GetValidatedObject<Model>(handle, outRealObj);
if (ret == RT_ERROR_NONE) {
return RT_ERROR_NONE;
} else {
return ReportApiHandleValidationError(ret, callerFuncName);
}
}
rtError_t ValidateLabelHandleForApi(rtLabel_t handle, Label*& outRealObj, const char_t* callerFuncName)
{
const rtError_t ret = GetValidatedObject<Label>(handle, outRealObj);
if (ret == RT_ERROR_NONE) {
return RT_ERROR_NONE;
} else {
return ReportApiHandleValidationError(ret, callerFuncName);
}
}
rtError_t ValidateLabelHandleArrayForApi(
rtLabel_t* handles, size_t count, std::vector<Label*>& outRealObjs, const char_t* callerFuncName)
{
outRealObjs.clear();
if (handles == nullptr) {
return ReportApiHandleValidationError(RT_ERROR_INVALID_VALUE, callerFuncName);
}
outRealObjs.reserve(count);
for (size_t i = 0U; i < count; ++i) {
Label* realObj = nullptr;
const rtError_t ret = GetValidatedObject<Label>(handles[i], realObj);
if (ret != RT_ERROR_NONE) {
return ReportApiHandleValidationError(ret, callerFuncName);
}
if (realObj == nullptr) {
return ReportApiHandleValidationError(RT_ERROR_INVALID_VALUE, callerFuncName);
}
outRealObjs.push_back(realObj);
}
return RT_ERROR_NONE;
}
rtError_t ValidateStreamHandleForApi(rtStream_t handle, Stream*& outRealObj, const char_t* callerFuncName)
{
const rtError_t ret = GetValidatedObject<Stream>(handle, outRealObj);
if (ret == RT_ERROR_NONE) {
return RT_ERROR_NONE;
} else {
return ReportApiHandleValidationError(ret, callerFuncName);
}
}
rtError_t ValidateEventHandleForApi(rtEvent_t handle, Event*& outRealObj, const char_t* callerFuncName)
{
const rtError_t ret = GetValidatedObject<Event>(handle, outRealObj);
if (ret == RT_ERROR_NONE) {
return RT_ERROR_NONE;
} else {
return ReportApiHandleValidationError(ret, callerFuncName);
}
}
rtError_t ValidateNotifyHandleForApi(rtNotify_t handle, Notify*& outRealObj, const char_t* callerFuncName)
{
const rtError_t ret = GetValidatedObject<Notify>(handle, outRealObj);
if (ret == RT_ERROR_NONE) {
return RT_ERROR_NONE;
} else {
return ReportApiHandleValidationError(ret, callerFuncName);
}
}
rtError_t ValidateCountNotifyHandleForApi(rtCntNotify_t handle, CountNotify*& outRealObj, const char_t* callerFuncName)
{
const rtError_t ret = GetValidatedObject<CountNotify>(handle, outRealObj);
if (ret == RT_ERROR_NONE) {
return RT_ERROR_NONE;
} else {
return ReportApiHandleValidationError(ret, callerFuncName);
}
}
rtError_t ValidateCondHandleHandleForApi(rtCondHandle_t handle, CondHandle*& outRealObj, const char_t* callerFuncName)
{
const rtError_t ret = GetValidatedObject<CondHandle>(handle, outRealObj);
if (ret == RT_ERROR_NONE) {
return RT_ERROR_NONE;
} else {
return ReportApiHandleValidationError(ret, callerFuncName);
}
}
rtError_t ValidateProgramHandleForApi(rtBinHandle handle, Program*& outRealObj, const char_t* callerFuncName)
{
const rtError_t ret = GetValidatedObject<Program>(handle, outRealObj);
if (ret == RT_ERROR_NONE) {
return RT_ERROR_NONE;
} else {
return ReportApiHandleValidationError(ret, callerFuncName);
}
}
rtError_t ValidateKernelHandleForApi(const void* handle, Kernel*& outRealObj, const char_t* callerFuncName)
{
const rtError_t ret = GetValidatedObject<Kernel>(handle, outRealObj);
if (ret == RT_ERROR_NONE) {
return RT_ERROR_NONE;
} else {
return ReportApiHandleValidationError(ret, callerFuncName);
}
}
rtError_t ConvertFuncToKernel(
Api* const apiInstance, const void* func, Kernel*& kernel, const char_t* const callerFuncName)
{
return ConvertFuncToKernelImpl(apiInstance, func, kernel, callerFuncName);
}
rtError_t ConvertFuncToKernel(
ApiKernelFunc* const apiInstance, const void* func, Kernel*& kernel, const char_t* const callerFuncName)
{
return ConvertFuncToKernelImpl(apiInstance, func, kernel, callerFuncName);
}
rtError_t ValidateArgsHandleForApi(rtArgsHandle handle, RtArgsHandle*& outRealObj, const char_t* callerFuncName)
{
const rtError_t ret = ProbeValidatedObject<RtArgsHandle>(handle, outRealObj);
if (ret == RT_ERROR_NONE) {
return RT_ERROR_NONE;
}
RtArgsHandle* const legacyArgsHandle = RtPtrToPtr<RtArgsHandle*>(handle);
if (legacyArgsHandle != nullptr) {
RtArgsHandle* realArgsHandle = nullptr;
const rtError_t legacyRet = ProbeValidatedObject<RtArgsHandle>(
RtPtrToPtr<rtArgsHandle>(RtInnerHandleAccessor<RtArgsHandle>::Get(legacyArgsHandle)), realArgsHandle);
if ((legacyRet == RT_ERROR_NONE) && (realArgsHandle == legacyArgsHandle)) {
outRealObj = legacyArgsHandle;
return RT_ERROR_NONE;
}
}
return ReportApiHandleValidationError(ret, callerFuncName);
}
rtError_t ValidateArgsHandleForUserMemApi(rtArgsHandle handle, RtArgsHandle*& outRealObj, const char_t* callerFuncName)
{
if (handle == nullptr) {
outRealObj = nullptr;
return RT_ERROR_NONE;
}
if (HasExpectedMagic(handle, RT_ARGS_HANDLE_MAGIC)) {
return ValidateArgsHandleForApi(handle, outRealObj, callerFuncName);
}
RtArgsHandle* const legacyArgsHandle = RtPtrToPtr<RtArgsHandle*>(handle);
RtArgsHandle* realArgsHandle = nullptr;
const rtError_t ret = GetValidatedObject<RtArgsHandle>(
RtPtrToPtr<rtArgsHandle>(RtInnerHandleAccessor<RtArgsHandle>::Get(legacyArgsHandle)), realArgsHandle);
if ((ret == RT_ERROR_NONE) && (realArgsHandle == legacyArgsHandle)) {
outRealObj = legacyArgsHandle;
return RT_ERROR_NONE;
}
return ReportApiHandleValidationError(RT_ERROR_INVALID_HANDLE, callerFuncName);
}
rtError_t ValidateParamHandleForApi(rtParaHandle handle, ParaDetail*& outRealObj, const char_t* callerFuncName)
{
const rtError_t ret = GetValidatedObject<ParaDetail>(handle, outRealObj);
if (ret == RT_ERROR_NONE) {
return RT_ERROR_NONE;
} else {
return ReportApiHandleValidationError(ret, callerFuncName);
}
}
rtError_t ValidateLaunchArgsHandleForApi(
rtLaunchArgsHandle handle, rtLaunchArgs_t*& outRealObj, const char_t* callerFuncName)
{
const rtError_t ret = ProbeValidatedObject<rtLaunchArgs_t>(handle, outRealObj);
if (ret == RT_ERROR_NONE) {
return RT_ERROR_NONE;
}
rtLaunchArgs_t* const legacyLaunchArgs = RtPtrToPtr<rtLaunchArgs_t*>(handle);
if (legacyLaunchArgs != nullptr) {
rtLaunchArgs_t* realLaunchArgs = nullptr;
const rtError_t legacyRet = ProbeValidatedObject<rtLaunchArgs_t>(
RtPtrToPtr<rtLaunchArgsHandle>(RtInnerHandleAccessor<rtLaunchArgs_t>::Get(legacyLaunchArgs)),
realLaunchArgs);
if ((legacyRet == RT_ERROR_NONE) && (realLaunchArgs == legacyLaunchArgs)) {
outRealObj = legacyLaunchArgs;
return RT_ERROR_NONE;
}
}
return ReportApiHandleValidationError(ret, callerFuncName);
}
}
}