Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
37 changes: 29 additions & 8 deletions source/api_c/src/c_api.cc
Original file line number Diff line number Diff line change
Expand Up @@ -730,6 +730,27 @@ inline void flatten_vector(std::vector<VALUETYPE>& onedv,
}
}

namespace {

constexpr char model_devi_nframes_error[] =
"DeePMD-kit Error: Model-deviation C APIs support exactly one frame.";

bool validate_model_devi_nframes(DP_DeepBaseModelDevi* dp, const int nframes) {
if (nframes == 1) {
return true;
}

// These helpers serve extern "C" entry points, so an unsupported frame
// count must be reported through DP_*CheckOK instead of unwinding a C++
// exception into a C caller. Keep this validation before constructing any
// pointer ranges or accessing the model/neighbor list: invalid input must
// remain safe even when those arguments are otherwise unusable.
dp->exception = model_devi_nframes_error;
Comment thread
njzjz marked this conversation as resolved.
return false;
}

} // namespace

template <typename VALUETYPE>
void DP_DeepPotModelDeviCompute_variant(
DP_DeepPotModelDevi* dp,
Expand All @@ -746,8 +767,8 @@ void DP_DeepPotModelDeviCompute_variant(
VALUETYPE* atomic_energy,
VALUETYPE* atomic_virial,
const VALUETYPE* charge_spin = nullptr) {
if (nframes > 1) {
throw std::runtime_error("nframes > 1 not supported yet");
if (!validate_model_devi_nframes(dp, nframes)) {
return;
}
// init C++ vectors from C arrays
std::vector<VALUETYPE> coord_(coord, coord + natoms * 3);
Expand Down Expand Up @@ -857,8 +878,8 @@ void DP_DeepSpinModelDeviCompute_variant(DP_DeepSpinModelDevi* dp,
VALUETYPE* virial,
VALUETYPE* atomic_energy,
VALUETYPE* atomic_virial) {
if (nframes > 1) {
throw std::runtime_error("nframes > 1 not supported yet");
if (!validate_model_devi_nframes(dp, nframes)) {
return;
}
// init C++ vectors from C arrays
std::vector<VALUETYPE> coord_(coord, coord + natoms * 3);
Expand Down Expand Up @@ -972,8 +993,8 @@ void DP_DeepPotModelDeviComputeNList_variant(
VALUETYPE* atomic_energy,
VALUETYPE* atomic_virial,
const VALUETYPE* charge_spin = nullptr) {
if (nframes > 1) {
throw std::runtime_error("nframes > 1 not supported yet");
if (!validate_model_devi_nframes(dp, nframes)) {
return;
}
// init C++ vectors from C arrays
std::vector<VALUETYPE> coord_(coord, coord + natoms * 3);
Expand Down Expand Up @@ -1097,8 +1118,8 @@ void DP_DeepSpinModelDeviComputeNList_variant(DP_DeepSpinModelDevi* dp,
VALUETYPE* virial,
VALUETYPE* atomic_energy,
VALUETYPE* atomic_virial) {
if (nframes > 1) {
throw std::runtime_error("nframes > 1 not supported yet");
if (!validate_model_devi_nframes(dp, nframes)) {
return;
}
// init C++ vectors from C arrays
std::vector<VALUETYPE> coord_(coord, coord + natoms * 3);
Expand Down
165 changes: 165 additions & 0 deletions source/api_c/tests/test_deepmd_exception.cc
Original file line number Diff line number Diff line change
Expand Up @@ -11,8 +11,32 @@
#include <utility>
#include <vector>

#include "c_api.h"
#include "c_api_internal.h"
#include "deepmd.hpp"

namespace {

constexpr char model_devi_nframes_error[] =
"DeePMD-kit Error: Model-deviation C APIs support exactly one frame.";

template <typename MODEL, typename INVOKE, typename CHECK_OK>
void expect_model_devi_frame_error(const INVOKE& invoke,
const CHECK_OK& check_ok) {
// Use a fresh error carrier for every public entry point. Reusing a model
// would let a previous failure remain visible through CheckOK and could hide
// a wrapper that forgot to validate its own frame count.
MODEL model;
ASSERT_NO_THROW(invoke(&model));

const char* error = check_ok(&model);
ASSERT_NE(error, nullptr);
EXPECT_STREQ(error, model_devi_nframes_error);
DP_DeleteChar(error);
}

} // namespace

TEST(TestDeepmdException, deepmdexception) {
std::string expected_error_message = "DeePMD-kit C API Error: unittest";
try {
Expand All @@ -27,6 +51,147 @@ TEST(TestDeepmdException, deepmdexception_nofile) {
deepmd::hpp::deepmd_exception);
}

TEST(TestModelDeviCAPIExceptionBoundary,
rejects_two_frames_across_all_multiframe_public_variants) {
constexpr int nframes = 2;
constexpr int natoms = 1;
const int atype[natoms] = {0};
const double coord[natoms * 3] = {0.0, 0.0, 0.0};
const float coordf[natoms * 3] = {0.0F, 0.0F, 0.0F};
const double spin[natoms * 3] = {0.0, 0.0, 0.0};
const float spinf[natoms * 3] = {0.0F, 0.0F, 0.0F};
const double charge_spin[1] = {0.0};
const float charge_spinf[1] = {0.0F};
DP_Nlist nlist;

const auto check_pot = [](DP_DeepPotModelDevi* model) {
return DP_DeepPotModelDeviCheckOK(model);
};
const auto check_spin = [](DP_DeepSpinModelDevi* model) {
return DP_DeepSpinModelDeviCheckOK(model);
};

// The four legacy model-deviation entry points have no nframes argument and
// always delegate with one frame. The twelve calls below cover every public
// entry point through which a caller can supply an unsupported frame count.
// Output and optional-parameter pointers may be null by contract. A default
// model and neighbor list deliberately keep this regression free of
// backend/model fixtures while verifying validation precedes model access.
expect_model_devi_frame_error<DP_DeepPotModelDevi>(
[&](DP_DeepPotModelDevi* model) {
DP_DeepPotModelDeviCompute2(model, nframes, natoms, coord, atype,
nullptr, nullptr, nullptr, nullptr, nullptr,
nullptr, nullptr, nullptr);
},
check_pot);
expect_model_devi_frame_error<DP_DeepPotModelDevi>(
[&](DP_DeepPotModelDevi* model) {
DP_DeepPotModelDeviComputef2(model, nframes, natoms, coordf, atype,
nullptr, nullptr, nullptr, nullptr,
nullptr, nullptr, nullptr, nullptr);
},
check_pot);
expect_model_devi_frame_error<DP_DeepSpinModelDevi>(
[&](DP_DeepSpinModelDevi* model) {
DP_DeepSpinModelDeviCompute2(
model, nframes, natoms, coord, spin, atype, nullptr, nullptr,
nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr);
},
check_spin);
expect_model_devi_frame_error<DP_DeepSpinModelDevi>(
[&](DP_DeepSpinModelDevi* model) {
DP_DeepSpinModelDeviComputef2(
model, nframes, natoms, coordf, spinf, atype, nullptr, nullptr,
nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr);
},
check_spin);

expect_model_devi_frame_error<DP_DeepPotModelDevi>(
[&](DP_DeepPotModelDevi* model) {
DP_DeepPotModelDeviComputeNList2(
model, nframes, natoms, coord, atype, nullptr, 0, &nlist, 0,
nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr);
},
check_pot);
expect_model_devi_frame_error<DP_DeepPotModelDevi>(
[&](DP_DeepPotModelDevi* model) {
DP_DeepPotModelDeviComputeNListf2(
model, nframes, natoms, coordf, atype, nullptr, 0, &nlist, 0,
nullptr, nullptr, nullptr, nullptr, nullptr, nullptr, nullptr);
},
check_pot);
expect_model_devi_frame_error<DP_DeepSpinModelDevi>(
[&](DP_DeepSpinModelDevi* model) {
DP_DeepSpinModelDeviComputeNList2(model, nframes, natoms, coord, spin,
atype, nullptr, 0, &nlist, 0, nullptr,
nullptr, nullptr, nullptr, nullptr,
nullptr, nullptr, nullptr);
},
check_spin);
expect_model_devi_frame_error<DP_DeepSpinModelDevi>(
[&](DP_DeepSpinModelDevi* model) {
DP_DeepSpinModelDeviComputeNListf2(model, nframes, natoms, coordf,
spinf, atype, nullptr, 0, &nlist, 0,
nullptr, nullptr, nullptr, nullptr,
nullptr, nullptr, nullptr, nullptr);
},
check_spin);

expect_model_devi_frame_error<DP_DeepPotModelDevi>(
[&](DP_DeepPotModelDevi* model) {
DP_DeepPotModelDeviCompute3(
model, nframes, natoms, coord, atype, nullptr, nullptr, nullptr,
charge_spin, nullptr, nullptr, nullptr, nullptr, nullptr);
},
check_pot);
expect_model_devi_frame_error<DP_DeepPotModelDevi>(
[&](DP_DeepPotModelDevi* model) {
DP_DeepPotModelDeviComputef3(
model, nframes, natoms, coordf, atype, nullptr, nullptr, nullptr,
charge_spinf, nullptr, nullptr, nullptr, nullptr, nullptr);
},
check_pot);
expect_model_devi_frame_error<DP_DeepPotModelDevi>(
[&](DP_DeepPotModelDevi* model) {
DP_DeepPotModelDeviComputeNList3(model, nframes, natoms, coord, atype,
nullptr, 0, &nlist, 0, nullptr,
nullptr, charge_spin, nullptr, nullptr,
nullptr, nullptr, nullptr);
},
check_pot);
expect_model_devi_frame_error<DP_DeepPotModelDevi>(
[&](DP_DeepPotModelDevi* model) {
DP_DeepPotModelDeviComputeNListf3(model, nframes, natoms, coordf, atype,
nullptr, 0, &nlist, 0, nullptr,
nullptr, charge_spinf, nullptr,
nullptr, nullptr, nullptr, nullptr);
},
check_pot);
}

class TestModelDeviInvalidFrameCount : public ::testing::TestWithParam<int> {};

TEST_P(TestModelDeviInvalidFrameCount, rejects_every_count_except_one) {
constexpr int natoms = 1;

expect_model_devi_frame_error<DP_DeepPotModelDevi>(
[&](DP_DeepPotModelDevi* model) {
// Required array pointers are deliberately unusable. The unsupported
// frame count must be rejected before pointer-range construction;
// otherwise zero/negative nframes could reach undefined pointer math.
DP_DeepPotModelDeviCompute2(model, GetParam(), natoms, nullptr, nullptr,
nullptr, nullptr, nullptr, nullptr, nullptr,
nullptr, nullptr, nullptr);
},
[](DP_DeepPotModelDevi* model) {
return DP_DeepPotModelDeviCheckOK(model);
});
}

INSTANTIATE_TEST_SUITE_P(UnsupportedFrameCounts,
TestModelDeviInvalidFrameCount,
::testing::Values(-1, 0, 2));

TEST(TestInputNlist, move_assignment_transfers_c_handle) {
deepmd::hpp::InputNlist source;
DP_Nlist* source_handle = source.nl;
Expand Down
Loading