diff --git a/tests/python/test_construct_chat_template_validation.py b/tests/python/test_construct_chat_template_validation.py index 66d3d80920..53d281d435 100644 --- a/tests/python/test_construct_chat_template_validation.py +++ b/tests/python/test_construct_chat_template_validation.py @@ -104,3 +104,93 @@ def test_chat_template_does_not_leak_sentinel_when_section_starts_with_it(chat_t ) assert "{INPUT}" not in jinja_template assert "{OUTPUT}" not in jinja_template + + +_SYSTEM_CHAT_TEMPLATE = ( + "{SYSTEM}\n" + "### User: {INPUT}\n### Assistant: {OUTPUT}" + "### User: {INPUT}\n### Assistant: {OUTPUT}" +) + + +def _render(jinja_template, messages): + from jinja2.sandbox import ImmutableSandboxedEnvironment + + env = ImmutableSandboxedEnvironment() + env.globals["raise_exception"] = lambda message: (_ for _ in ()).throw(RuntimeError(message)) + return env.from_string(jinja_template).render( + messages = messages, + bos_token = "", + eos_token = "", + add_generation_prompt = False, + ) + + +@pytest.mark.parametrize("default_system_message", [None, "You are helpful."]) +def test_system_message_is_consumed_by_the_system_part(default_system_message): + """A caller-supplied system message must be rendered by the system part and + skipped by the message loop, whatever `default_system_message` is. + + With `default_system_message = None` the generated template used to bind + `loop_messages` only inside the `{% if %}` arm. The `Fix missing + loop_messages` step then saw no unconditional binding, rewrote the loop back + to `messages`, and the system message reached the loop and tripped + `raise_exception`. + """ + _, jinja_template, _, _ = construct_chat_template( + tokenizer = _SuccessFakeTokenizer(), + chat_template = _SYSTEM_CHAT_TEMPLATE, + default_system_message = default_system_message, + extra_eos_tokens = [""], + ) + rendered = _render( + jinja_template, + [ + {"role": "system", "content": "Be terse."}, + {"role": "user", "content": "Hi"}, + ], + ) + assert rendered.count("Be terse.") == 1 + assert rendered.count("Hi") == 1 + # A caller system message overrides the default; the default must not leak in. + if default_system_message is not None: + assert default_system_message not in rendered + + +def test_absent_system_message_still_renders_without_default(): + """`default_system_message = None` with no system message in the input must + keep working -- the `{% else %}` arm has to bind `loop_messages = messages`.""" + _, jinja_template, _, _ = construct_chat_template( + tokenizer = _SuccessFakeTokenizer(), + chat_template = _SYSTEM_CHAT_TEMPLATE, + default_system_message = None, + extra_eos_tokens = [""], + ) + rendered = _render(jinja_template, [{"role": "user", "content": "Hi"}]) + assert "Hi" in rendered + + +_NO_SYSTEM_CHAT_TEMPLATE = ( + "PREAMBLE\n" + "### User: {INPUT}\n### Assistant: {OUTPUT}" + "### User: {INPUT}\n### Assistant: {OUTPUT}" +) + + +def test_static_prefix_without_system_still_rejects_system_message(): + """A template with a static prefix but no {SYSTEM} placeholder cannot render a + caller system message, so it must still raise rather than silently drop it.""" + _, jinja_template, _, _ = construct_chat_template( + tokenizer = _SuccessFakeTokenizer(), + chat_template = _NO_SYSTEM_CHAT_TEMPLATE, + default_system_message = None, + extra_eos_tokens = [""], + ) + with pytest.raises(RuntimeError, match = "Only user and assistant roles are supported!"): + _render( + jinja_template, + [ + {"role": "system", "content": "Be terse."}, + {"role": "user", "content": "Hi"}, + ], + ) diff --git a/unsloth/chat_templates.py b/unsloth/chat_templates.py index f47c78ba80..b857c34bcb 100644 --- a/unsloth/chat_templates.py +++ b/unsloth/chat_templates.py @@ -2652,6 +2652,12 @@ extra_eos_tokens = None, "{{ '" + full_system + "' }}"\ "{% set loop_messages = messages %}"\ "{% endif %}" + elif "{SYSTEM}" in system_part: + # Only bind loop_messages when the template can render a caller system + # message. A static prefix with no {SYSTEM} must still raise, not drop it. + partial_system += "{% else %}"\ + "{% set loop_messages = messages %}"\ + "{% endif %}" else: partial_system += "{% endif %}"