unsloth/tests/python/test_change_system_message.py
Daniel Han 44f113cf0d
Escape the system message spliced into predefined chat templates (#7746)
* Escape the system message spliced into predefined chat templates

get_chat_template(..., system_message = ...) substitutes the message into a
{system_message} placeholder that sits inside a Jinja string literal in all 15
predefined templates that carry one, so a quote closes the literal and a
backslash is read as an escape:

  vicuna  "Answer the user's question."  -> TemplateSyntaxError
  vicuna  r"Put it in \boxed{}."         -> renders '\x08oxed{}'
  vicuna  r"C:\Users\me"                 -> TemplateSyntaxError

Reuse the escaper PR #7731 added for construct_chat_template, promoted to a
module-level _escape_jinja_literal and extended to escape double quotes so the
one helper covers llama-3.1's "..." literal as well as the '...' the rest use.
Apply it to the predefined branch of _change_system_message and to the ShareGPT
mapping values, and drop the hand-escaping from the two vicuna defaults, which
would otherwise be escaped twice.

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Tighten the escaping comments

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
2026-08-02 07:18:09 -07:00

82 lines
3.2 KiB
Python

import ast
import re
import types
from pathlib import Path
import pytest
def _load_change_system_message():
# Extract _change_system_message without importing unsloth (needs unsloth_zoo / a GPU).
source = Path(__file__).parents[2] / "unsloth" / "chat_templates.py"
tree = ast.parse(source.read_text(encoding = "utf-8"))
funcs = [
node
for node in tree.body
if isinstance(node, ast.FunctionDef)
and node.name in ("_change_system_message", "_escape_jinja_literal")
]
namespace = {
"re": re,
"logger": types.SimpleNamespace(warning_once = lambda *a, **k: None),
"DEFAULT_SYSTEM_MESSAGE": {"unsloth": "You are a helpful assistant to the user"},
}
module = ast.Module(body = funcs, type_ignores = [])
ast.fix_missing_locations(module)
exec(compile(module, str(source), "exec"), namespace)
return namespace["_change_system_message"]
CUSTOM = "mycustom" # no predefined default
def test_custom_template_fills_placeholder():
# A {system_message} placeholder must be filled, not left literal.
fn = _load_change_system_message()
template, used = fn("System: {system_message}\nUser:", CUSTOM, "You are a pirate")
assert template == "System: You are a pirate\nUser:"
assert "{system_message}" not in template
assert used == "You are a pirate"
def test_custom_template_preserves_backslashes():
# str.replace not re.sub: re.sub treats backslashes specially (r"C:\Users"
# bad-escape, r"\1" group ref), so messages must be inserted verbatim.
fn = _load_change_system_message()
for msg in (r"C:\Users\me", r"\frac{a}{b}", r"see \1 here"):
template, used = fn("System: {system_message}", CUSTOM, msg)
assert template == f"System: {msg}"
assert used == msg
def test_custom_template_requires_system_message():
# A placeholder with no system message must raise, not stay literal.
fn = _load_change_system_message()
with pytest.raises(ValueError):
fn("System: {system_message}", CUSTOM, None)
def test_custom_template_without_placeholder_unchanged():
fn = _load_change_system_message()
template, used = fn("System: fixed", CUSTOM, "ignored")
assert template == "System: fixed"
def test_predefined_template_uses_default_then_override():
fn = _load_change_system_message()
t1, u1 = fn("System: {system_message}", "unsloth", None)
assert t1 == "System: You are a helpful assistant to the user"
t2, u2 = fn("System: {system_message}", "unsloth", "Custom override")
assert t2 == "System: Custom override"
assert u2 == "Custom override"
def test_predefined_template_escapes_but_custom_does_not():
# Predefined templates hold {system_message} inside a Jinja literal, so it is escaped;
# a custom template's placeholder may be raw text, so it stays verbatim.
fn = _load_change_system_message()
msg = """it's a \\test "x"."""
predefined, used = fn("System: {system_message}", "unsloth", msg)
assert predefined == """System: it\\'s a \\\\test \\"x\\"."""
assert used == msg # returned message is the raw one, not the escaped source
assert fn("System: {system_message}", CUSTOM, msg)[0] == f"System: {msg}"