| 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": {}}], |
| } |
| ] |
| }, |
| ) |
|
|