diff --git a/src/models/model.cpp b/src/models/model.cpp index 06e0dda316..444be44529 100644 --- a/src/models/model.cpp +++ b/src/models/model.cpp @@ -335,6 +335,15 @@ std::vector Tokenizer::EncodeBatch(std::span strings } std::shared_ptr Tokenizer::EncodeBatch(std::span strings) const { + if (strings.empty()) { + throw std::runtime_error("EncodeBatch: input strings must not be empty"); + } + for (size_t i = 0; i < strings.size(); i++) { + if (strings[i] == nullptr) { + throw std::runtime_error("EncodeBatch: input string at index " + std::to_string(i) + " must not be null"); + } + } + std::vector> sequences; std::vector> span_sequences; for (size_t i = 0; i < strings.size(); i++) { diff --git a/src/ort_genai_c.cpp b/src/ort_genai_c.cpp index 3ea17eaca7..cfb8966b3d 100644 --- a/src/ort_genai_c.cpp +++ b/src/ort_genai_c.cpp @@ -662,6 +662,8 @@ OgaResult* OGA_API_CALL OgaTokenizerEncode(const OgaTokenizer* tokenizer, const OgaResult* OGA_API_CALL OgaTokenizerEncodeBatch(const OgaTokenizer* tokenizer, const char** strings, size_t count, OgaTensor** out) { OGA_TRY + if (count > 0 && strings == nullptr) + throw std::runtime_error("EncodeBatch: strings pointer must not be null when count > 0"); auto tensor = tokenizer->EncodeBatch(std::span(strings, count)); *out = ReturnShared(tensor); return nullptr; diff --git a/test/c_api_tests.cpp b/test/c_api_tests.cpp index f6e952a9bd..e2a5135a29 100644 --- a/test/c_api_tests.cpp +++ b/test/c_api_tests.cpp @@ -133,6 +133,21 @@ TEST(CAPITests, TokenizerCAPI) { #endif } +TEST(CAPITests, EncodeBatchEmptyInputThrows) { +#if TEST_PHI2 + auto model = OgaModel::Create(PHI2_PATH); + auto tokenizer = OgaTokenizer::Create(*model); + + // EncodeBatch with zero strings should throw, not crash with SIGFPE + ASSERT_THROW(tokenizer->EncodeBatch(nullptr, 0), std::runtime_error); + + // Invalid pointers with count > 0 should also be rejected deterministically. + ASSERT_THROW(tokenizer->EncodeBatch(nullptr, 1), std::runtime_error); + const char* bad_strings[] = {nullptr}; + ASSERT_THROW(tokenizer->EncodeBatch(bad_strings, 1), std::runtime_error); +#endif +} + TEST(CAPITests, TokenizerUpdateOptions) { #if TEST_PHI2 auto config = OgaConfig::Create(PHI2_PATH);