Skip to content
Open
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
55 changes: 55 additions & 0 deletions shared/api/minja.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<char>(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<char>(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<std::string>();
auto after = vargs.args[1].get<std::string>();
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<int64_t>::max)();
if (vargs.args.size() == 3)
{
auto requested = vargs.args[2].get<int64_t>();
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);
Comment thread
Copilot marked this conversation as resolved.
}
}
throw std::runtime_error("Unknown method: " + method->get_name());
}
Expand Down
88 changes: 88 additions & 0 deletions test/pp_api_test/test_tokenizer_chat.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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: <name>". 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<OrtxTensorResult> 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<OrtxTensor> 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<const void**>(&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<OrtxTokenizer> 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<OrtxTokenizer> 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");
}
Loading