custom-chat-template / tests /test_chat_template.py
Berk Birkan
Add custom ChatML tool-calling template
6a438b2
Raw
History Blame Contribute Delete
3.31 kB
import json
import re
from pathlib import Path
import jinja2
import pytest
ROOT = Path(__file__).resolve().parents[1]
def raise_exception(message: str) -> None:
raise jinja2.TemplateError(message)
@pytest.fixture(scope="module")
def template() -> jinja2.Template:
environment = jinja2.Environment(
undefined=jinja2.StrictUndefined,
trim_blocks=False,
lstrip_blocks=False,
)
environment.globals["raise_exception"] = raise_exception
return environment.from_string(
(ROOT / "chat_template.jinja").read_text(encoding="utf-8")
)
def load_example(name: str) -> dict:
return json.loads((ROOT / "examples" / name).read_text(encoding="utf-8"))
def render(template: jinja2.Template, payload: dict) -> str:
return template.render(
bos_token="",
messages=payload["messages"],
tools=payload.get("tools"),
add_generation_prompt=payload.get("add_generation_prompt", False),
)
def test_basic_roles_are_serialized(template: jinja2.Template) -> None:
output = render(template, load_example("basic_chat.json"))
assert "<|im_start|>system\n" in output
assert "<|im_start|>user\n" in output
assert "<|im_start|>assistant\n" in output
assert output.count("<|im_end|>") == 3
def test_tools_calls_and_responses_are_serialized(template: jinja2.Template) -> None:
output = render(template, load_example("tool_call_chat.json"))
assert "<tools>\n[" in output
assert '"type": "function"' in output
assert "<tool_call>\n" in output
assert '"name": "get_weather"' in output
tool_call_match = re.search(r"<tool_call>\n(.+?)\n</tool_call>", output)
assert tool_call_match is not None
assert json.loads(tool_call_match.group(1)) == {
"name": "get_weather",
"arguments": {"city": "İstanbul"},
}
assert "<tool_response>\n" in output
assert '"temperature_c":24' in output
def test_tools_get_a_default_system_message(template: jinja2.Template) -> None:
output = render(
template,
{
"tools": load_example("tool_call_chat.json")["tools"],
"messages": [{"role": "user", "content": "Hava nasıl?"}],
},
)
assert output.startswith("<|im_start|>system\n")
assert "Never invent tool results." in output
def test_generation_prompt_is_appended(template: jinja2.Template) -> None:
output = render(
template,
{
"messages": [{"role": "user", "content": "Merhaba"}],
"add_generation_prompt": True,
},
)
assert output.endswith("<|im_start|>assistant\n")
def test_unknown_role_fails_explicitly(template: jinja2.Template) -> None:
with pytest.raises(jinja2.TemplateError, match="Unsupported message role"):
render(template, {"messages": [{"role": "developer", "content": "x"}]})
def test_invalid_tool_call_fails_explicitly(template: jinja2.Template) -> None:
with pytest.raises(jinja2.TemplateError, match="function.name"):
render(
template,
{
"messages": [
{
"role": "assistant",
"tool_calls": [{"type": "function", "function": {}}],
}
]
},
)