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 "\n[" in output assert '"type": "function"' in output assert "\n" in output assert '"name": "get_weather"' in output tool_call_match = re.search(r"\n(.+?)\n", output) assert tool_call_match is not None assert json.loads(tool_call_match.group(1)) == { "name": "get_weather", "arguments": {"city": "İstanbul"}, } assert "\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": {}}], } ] }, )