File size: 4,249 Bytes
dfb775d
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
"""Chat templates — Hermes / Qwen3-Coder / Qwen3 reasoning parsers.

Pure-Python rendering and response parsing. Used by:
    1. `mindxtrain serve` to set the right `--chat-template` on vLLM-ROCm.
    2. `mindxtrain.operator.app` to format ChatRequest messages before forwarding.

Single canonical home per mindxtrain2.md §Part 4 `models.chat_template`. Merges
the previous `xtrain.serve.parsers` and `automindx.templates.registry` modules.
"""

from __future__ import annotations

import re
from collections.abc import Iterable
from dataclasses import dataclass
from typing import Literal, Protocol

Role = Literal["system", "user", "assistant", "tool"]


@dataclass(frozen=True)
class ChatMessage:
    role: Role
    content: str


class ChatTemplate(Protocol):
    """Callable interface: list[ChatMessage] -> rendered prompt str."""

    name: str

    def render(self, messages: Iterable[ChatMessage], add_generation_prompt: bool = True) -> str: ...
    def parse_response(self, response: str) -> dict[str, str]: ...


# ---- Hermes (ChatML) -------------------------------------------------------


class HermesTemplate:
    """ChatML-flavored format used by Hermes-3 / Qwen / many open models."""

    name: str = "hermes"

    def render(self, messages: Iterable[ChatMessage], add_generation_prompt: bool = True) -> str:
        parts: list[str] = []
        for m in messages:
            parts.append(f"<|im_start|>{m.role}\n{m.content}<|im_end|>")
        rendered = "\n".join(parts)
        if add_generation_prompt:
            rendered += "\n<|im_start|>assistant\n"
        return rendered

    def parse_response(self, response: str) -> dict[str, str]:
        cleaned = response.split("<|im_end|>", 1)[0].rstrip()
        return {"content": cleaned}


# ---- Qwen3-Coder -----------------------------------------------------------


class Qwen3CoderTemplate:
    """Qwen3-Coder uses Hermes-style framing plus a `<tool_call>` JSON tag."""

    name: str = "qwen3_coder"

    _TOOL_CALL_RE = re.compile(r"<tool_call>(.*?)</tool_call>", re.DOTALL)

    def render(self, messages: Iterable[ChatMessage], add_generation_prompt: bool = True) -> str:
        return HermesTemplate().render(messages, add_generation_prompt=add_generation_prompt)

    def parse_response(self, response: str) -> dict[str, str]:
        cleaned = response.split("<|im_end|>", 1)[0].rstrip()
        tool_calls = self._TOOL_CALL_RE.findall(cleaned)
        content = self._TOOL_CALL_RE.sub("", cleaned).strip()
        out: dict[str, str] = {"content": content}
        if tool_calls:
            out["tool_call"] = tool_calls[0].strip()
        return out


# ---- Qwen3 reasoning -------------------------------------------------------


class Qwen3ReasoningTemplate:
    """Qwen3 / Qwen3.5 / Qwen3.6 thinking format with `<think>...</think>` blocks."""

    name: str = "qwen3_reasoning"

    _THINK_RE = re.compile(r"<think>(.*?)</think>", re.DOTALL)

    def render(self, messages: Iterable[ChatMessage], add_generation_prompt: bool = True) -> str:
        return HermesTemplate().render(messages, add_generation_prompt=add_generation_prompt)

    def parse_response(self, response: str) -> dict[str, str]:
        cleaned = response.split("<|im_end|>", 1)[0].rstrip()
        thoughts = self._THINK_RE.findall(cleaned)
        content = self._THINK_RE.sub("", cleaned).strip()
        out: dict[str, str] = {"content": content}
        if thoughts:
            out["thinking"] = thoughts[0].strip()
        return out


# ---- registry --------------------------------------------------------------

_TEMPLATES: dict[str, ChatTemplate] = {
    "hermes": HermesTemplate(),
    "qwen3_coder": Qwen3CoderTemplate(),
    "qwen3": Qwen3ReasoningTemplate(),
    "qwen3_reasoning": Qwen3ReasoningTemplate(),
    "deepseek_r1": Qwen3ReasoningTemplate(),
}


def get_template(name: str) -> ChatTemplate:
    """Return the named template; default to Hermes if unknown."""
    return _TEMPLATES.get(name, _TEMPLATES["hermes"])


def list_templates() -> list[str]:
    """Return the names of all registered chat templates."""
    return sorted(_TEMPLATES)


# Back-compat alias for code that previously called `get_chat_template`.
get_chat_template = get_template