File size: 3,310 Bytes
6a438b2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
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": {}}],
                    }
                ]
            },
        )