studio_worker/engine/
chat_template.rs1use crate::types::ChatMessage;
9use serde_json::{Map, Value};
10
11#[derive(Debug, Clone, Default)]
13pub struct TemplateVars {
14 pub bos_token: String,
15 pub eos_token: String,
16}
17
18#[derive(Debug, thiserror::Error)]
20#[error("chat template: {0}")]
21pub struct TemplateError(String);
22
23pub fn render_chat(
26 template: &str,
27 messages: &[ChatMessage],
28 kwargs: &Map<String, Value>,
29 vars: &TemplateVars,
30) -> Result<String, TemplateError> {
31 let mut env = minijinja::Environment::new();
32 minijinja_contrib::add_to_environment(&mut env);
33 env.set_unknown_method_callback(minijinja_contrib::pycompat::unknown_method_callback);
34 env.add_function(
35 "raise_exception",
36 |message: String| -> Result<String, minijinja::Error> {
37 Err(minijinja::Error::new(
38 minijinja::ErrorKind::InvalidOperation,
39 message,
40 ))
41 },
42 );
43 let err = |e: minijinja::Error| {
44 let detail = e
46 .detail()
47 .map(str::to_string)
48 .unwrap_or_else(|| e.to_string());
49 TemplateError(detail)
50 };
51 let compiled = env.template_from_str(template).map_err(err)?;
52 let mut ctx: Map<String, Value> = kwargs.clone();
53 ctx.insert(
54 "messages".into(),
55 serde_json::to_value(messages).map_err(|e| TemplateError(e.to_string()))?,
56 );
57 ctx.insert("add_generation_prompt".into(), Value::Bool(true));
58 ctx.insert("bos_token".into(), vars.bos_token.clone().into());
59 ctx.insert("eos_token".into(), vars.eos_token.clone().into());
60 compiled
61 .render(minijinja::Value::from_serialize(&ctx))
62 .map_err(err)
63}
64
65pub fn merge_kwargs(
67 model: Option<&Map<String, Value>>,
68 request: Option<&Map<String, Value>>,
69) -> Map<String, Value> {
70 let mut merged = model.cloned().unwrap_or_default();
71 if let Some(request) = request {
72 merged.extend(request.clone());
73 }
74 merged
75}
76
77#[cfg(test)]
78mod tests {
79 use super::*;
80 use crate::types::ChatMessage;
81
82 const QWEN_LIKE: &str = r#"{%- for message in messages %}{{- '<|im_start|>' + message.role + '\n' + message.content + '<|im_end|>' + '\n' }}{%- endfor %}{%- if add_generation_prompt %}{{- '<|im_start|>assistant\n' }}{%- if enable_thinking is defined and enable_thinking is true %}{{- '<think>\n' }}{%- else %}{{- '<think>\n\n</think>\n\n' }}{%- endif %}{%- endif %}"#;
84
85 fn msgs() -> Vec<ChatMessage> {
86 vec![
87 ChatMessage {
88 role: "system".into(),
89 content: "Answer in JSON.".into(),
90 },
91 ChatMessage {
92 role: "user".into(),
93 content: "hi".into(),
94 },
95 ]
96 }
97
98 fn vars() -> TemplateVars {
99 TemplateVars {
100 bos_token: "<s>".into(),
101 eos_token: "</s>".into(),
102 }
103 }
104
105 fn kwargs(json: serde_json::Value) -> serde_json::Map<String, serde_json::Value> {
106 json.as_object().unwrap().clone()
107 }
108
109 #[test]
110 fn renders_messages_and_the_generation_prompt() {
111 let out = render_chat(QWEN_LIKE, &msgs(), &Default::default(), &vars()).unwrap();
112 assert_eq!(
113 out,
114 "<|im_start|>system\nAnswer in JSON.<|im_end|>\n<|im_start|>user\nhi<|im_end|>\n\
115 <|im_start|>assistant\n<think>\n\n</think>\n\n"
116 );
117 }
118
119 #[test]
120 fn kwargs_reach_the_template() {
121 let out = render_chat(
122 QWEN_LIKE,
123 &msgs(),
124 &kwargs(serde_json::json!({ "enable_thinking": true })),
125 &vars(),
126 )
127 .unwrap();
128 assert!(out.ends_with("<|im_start|>assistant\n<think>\n"), "{out}");
129 }
130
131 #[test]
132 fn bos_and_eos_tokens_are_available() {
133 let out = render_chat(
134 "{{ bos_token }}{{ messages[0].content }}{{ eos_token }}",
135 &msgs(),
136 &Default::default(),
137 &vars(),
138 )
139 .unwrap();
140 assert_eq!(out, "<s>Answer in JSON.</s>");
141 }
142
143 #[test]
144 fn python_string_methods_work() {
145 let out = render_chat(
146 "{% if messages[1].content.startswith('h') %}{{ messages[1].content.upper() }}{% endif %}",
147 &msgs(),
148 &Default::default(),
149 &vars(),
150 )
151 .unwrap();
152 assert_eq!(out, "HI");
153 }
154
155 #[test]
156 fn raise_exception_becomes_a_named_error() {
157 let err = render_chat(
158 "{{ raise_exception('System role not supported') }}",
159 &msgs(),
160 &Default::default(),
161 &vars(),
162 )
163 .unwrap_err();
164 assert!(
165 err.to_string().contains("System role not supported"),
166 "{err}"
167 );
168 }
169
170 #[test]
171 fn a_broken_template_is_a_named_error() {
172 let err = render_chat("{% for %}", &msgs(), &Default::default(), &vars()).unwrap_err();
173 assert!(err.to_string().starts_with("chat template"), "{err}");
174 }
175
176 const QWEN35: &str = include_str!("../../tests/fixtures/qwen3.5-chat-template.jinja");
179
180 #[test]
181 fn renders_the_real_qwen35_template_with_thinking_off() {
182 let out = render_chat(
183 QWEN35,
184 &msgs(),
185 &kwargs(serde_json::json!({ "enable_thinking": false })),
186 &vars(),
187 )
188 .unwrap();
189 assert!(
190 out.starts_with("<|im_start|>system\nAnswer in JSON.<|im_end|>\n"),
191 "{out}"
192 );
193 assert!(out.contains("<|im_start|>user\nhi<|im_end|>\n"), "{out}");
194 assert!(
195 out.ends_with("<|im_start|>assistant\n<think>\n\n</think>\n\n"),
196 "{out}"
197 );
198 }
199
200 #[test]
201 fn renders_the_real_qwen35_template_with_thinking_on() {
202 let out = render_chat(
203 QWEN35,
204 &msgs(),
205 &kwargs(serde_json::json!({ "enable_thinking": true })),
206 &vars(),
207 )
208 .unwrap();
209 assert!(out.ends_with("<|im_start|>assistant\n<think>\n"), "{out}");
210 }
211
212 #[test]
213 fn request_kwargs_override_model_defaults() {
214 let merged = merge_kwargs(
215 Some(&kwargs(
216 serde_json::json!({ "enable_thinking": false, "a": 1 }),
217 )),
218 Some(&kwargs(serde_json::json!({ "enable_thinking": true }))),
219 );
220 assert_eq!(
221 serde_json::Value::Object(merged),
222 serde_json::json!({ "enable_thinking": true, "a": 1 })
223 );
224 assert!(merge_kwargs(None, None).is_empty());
225 }
226}