diff --git a/shared/api/minja.hpp b/shared/api/minja.hpp index 7bf110942..5426ce2f7 100644 --- a/shared/api/minja.hpp +++ b/shared/api/minja.hpp @@ -2212,6 +2212,61 @@ namespace minja } return res; } + else if (method->get_name() == "upper") + { + vargs.expectArgs("upper method", {0, 0}, {0, 0}); + // ASCII-only case mapping for deterministic, locale-independent + // behavior. Matches upstream minja in practice for ASCII inputs; + // non-ASCII bytes are passed through unchanged. Full Unicode + // case folding is out of scope for the parser. + auto res = str; + for (char& c : res) { + if (c >= 'a' && c <= 'z') c = static_cast(c - ('a' - 'A')); + } + return Value(res); + } + else if (method->get_name() == "lower") + { + vargs.expectArgs("lower method", {0, 0}, {0, 0}); + // ASCII-only case mapping; see .upper() comment above. + auto res = str; + for (char& c : res) { + if (c >= 'A' && c <= 'Z') c = static_cast(c + ('a' - 'A')); + } + return Value(res); + } + else if (method->get_name() == "replace") + { + vargs.expectArgs("replace method", {2, 3}, {0, 0}); + auto before = vargs.args[0].get(); + auto after = vargs.args[1].get(); + auto res = str; + if (before.empty()) + { + return Value(res); + } + // Python str.replace semantics: count < 0 means "replace all"; + // count >= 0 limits the number of replacements; omitted argument + // also means "replace all". + int64_t count = (std::numeric_limits::max)(); + if (vargs.args.size() == 3) + { + auto requested = vargs.args[2].get(); + if (requested >= 0) + { + count = requested; + } + } + size_t start_pos = 0; + while (count > 0 && + (start_pos = res.find(before, start_pos)) != std::string::npos) + { + res.replace(start_pos, before.length(), after); + start_pos += after.length(); + --count; + } + return Value(res); + } } throw std::runtime_error("Unknown method: " + method->get_name()); } diff --git a/test/pp_api_test/test_tokenizer_chat.cc b/test/pp_api_test/test_tokenizer_chat.cc index 90e5f3f03..c1e7abe19 100644 --- a/test/pp_api_test/test_tokenizer_chat.cc +++ b/test/pp_api_test/test_tokenizer_chat.cc @@ -2255,4 +2255,92 @@ TEST(OrtxTokenizerTest, ChatTemplateDivisionByZero) { messages_json.c_str(), nullptr, result.ToBeAssigned(), false, false); EXPECT_NE(err, kOrtxOK) << "Expected modulo by zero to return an error."; } +} + +// Regression tests for missing string methods (.replace, .upper, .lower) in +// the vendored minja parser. Chat templates from popular HF models (e.g. +// smollm3-3b) call these on string values; previously they failed with +// "Unknown method: ". See: +// - microsoft/onnxruntime-extensions#1081 +// - microsoft/Foundry-Local#800 +namespace { + +void ExpectTemplateRenders(OrtxTokenizer* tokenizer, + const std::string& tmpl, + const std::string& expected) { + std::string messages_json = R"([{"role":"user","content":"hi"}])"; + OrtxObjectPtr result; + auto err = OrtxApplyChatTemplate( + tokenizer, tmpl.c_str(), messages_json.c_str(), nullptr, + result.ToBeAssigned(), false, false); + ASSERT_EQ(err, kOrtxOK) << "template: " << tmpl + << " err: " << OrtxGetLastErrorMessage(); + + OrtxObjectPtr tensor; + err = OrtxTensorResultGetAt(result.get(), 0, tensor.ToBeAssigned()); + ASSERT_EQ(err, kOrtxOK) << "OrtxTensorResultGetAt failed for template: " << tmpl; + ASSERT_EQ(tensor.Code(), kOrtxOK); + const char* text = nullptr; + err = OrtxGetTensorData(tensor.get(), reinterpret_cast(&text), + nullptr, nullptr); + ASSERT_EQ(err, kOrtxOK) << "OrtxGetTensorData failed for template: " << tmpl; + ASSERT_NE(text, nullptr) << "OrtxGetTensorData returned null data for template: " << tmpl; + EXPECT_STREQ(text, expected.c_str()) << "template: " << tmpl; +} + +} // namespace + +TEST(OrtxTokenizerTest, MinjaStringReplace) { + OrtxObjectPtr tokenizer(OrtxCreateTokenizer, "data/phi-4-base"); + ASSERT_EQ(tokenizer.Code(), kOrtxOK); + + // Single replacement. + ExpectTemplateRenders(tokenizer.get(), + "{{ 'foo and foo'.replace('foo', 'bar') }}", + "bar and bar"); + + // Chained replace + rstrip, exactly the pattern smollm3-3b uses. + ExpectTemplateRenders( + tokenizer.get(), + "{{ 'be helpful /no_think and /think '" + ".replace('/no_think', '').replace('/think', '').rstrip() }}", + "be helpful and"); + + // Replace with optional count argument (Python str.replace semantics). + ExpectTemplateRenders(tokenizer.get(), + "{{ 'aaaa'.replace('a', 'X', 2) }}", + "XXaa"); + + // Python str.replace semantics: count < 0 means "replace all". + ExpectTemplateRenders(tokenizer.get(), + "{{ 'aaaa'.replace('a', 'X', -1) }}", + "XXXX"); + + // Python str.replace semantics: count == 0 means no replacements. + ExpectTemplateRenders(tokenizer.get(), + "{{ 'aaaa'.replace('a', 'X', 0) }}", + "aaaa"); + + // Replace with empty `before` is a no-op (matches upstream minja behavior). + ExpectTemplateRenders(tokenizer.get(), + "{{ 'abc'.replace('', 'X') }}", + "abc"); + + // No match: original string returned unchanged. + ExpectTemplateRenders(tokenizer.get(), + "{{ 'abc'.replace('z', 'X') }}", + "abc"); +} + +TEST(OrtxTokenizerTest, MinjaStringUpperLower) { + OrtxObjectPtr tokenizer(OrtxCreateTokenizer, "data/phi-4-base"); + ASSERT_EQ(tokenizer.Code(), kOrtxOK); + + ExpectTemplateRenders(tokenizer.get(), "{{ 'Hello, World!'.upper() }}", + "HELLO, WORLD!"); + ExpectTemplateRenders(tokenizer.get(), "{{ 'Hello, World!'.lower() }}", + "hello, world!"); + // Idempotent on already-cased strings. + ExpectTemplateRenders(tokenizer.get(), "{{ 'ABC'.upper() }}", "ABC"); + ExpectTemplateRenders(tokenizer.get(), "{{ 'abc'.lower() }}", "abc"); } \ No newline at end of file