use minijinja::value::Value as JValue;
use minijinja::{Environment, Error, ErrorKind, State};
use serde_json::Value;
fn unknown_method(
_state: &State,
value: &JValue,
method: &str,
args: &[JValue],
) -> Result<JValue, Error> {
match method {
"get" => {
let key = args
.first()
.ok_or_else(|| Error::new(ErrorKind::MissingArgument, "get() requires a key"))?;
let default = args.get(1).cloned().unwrap_or_else(|| JValue::from(()));
match value.get_item(key) {
Ok(found) if !found.is_undefined() => Ok(found),
_ => Ok(default),
}
}
"split" => {
let text = value.as_str().ok_or_else(|| {
Error::new(ErrorKind::InvalidOperation, "split() called on a non-string")
})?;
let parts: Vec<JValue> = match args.first().and_then(|a| a.as_str()) {
Some(sep) => text.split(sep).map(JValue::from).collect(),
None => text.split_whitespace().map(JValue::from).collect(),
};
Ok(JValue::from(parts))
}
_ => Err(Error::from(ErrorKind::UnknownMethod)),
}
}
pub(crate) struct TemplateVars<'a> {
pub enable_thinking: bool,
pub bos_token: &'a str,
}
pub(crate) fn render_chat_template(
template: &str,
messages: &[Value],
vars: &TemplateVars<'_>,
) -> Result<String, Error> {
let mut env = Environment::new();
env.set_unknown_method_callback(unknown_method);
env.set_keep_trailing_newline(true);
env.add_template("chat", template)?;
let tmpl = env.get_template("chat")?;
tmpl.render(minijinja::context! {
messages => JValue::from_serialize(messages),
tools => JValue::from(()),
enable_thinking => vars.enable_thinking,
preserve_thinking => false,
add_generation_prompt => true,
bos_token => vars.bos_token,
eos_token => "",
})
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn msgs() -> Vec<Value> {
vec![
json!({"role": "system", "content": "Be brief."}),
json!({"role": "user", "content": "Hi"}),
]
}
fn vars<'a>(thinking: bool) -> TemplateVars<'a> {
TemplateVars {
enable_thinking: thinking,
bos_token: "",
}
}
#[test]
fn renders_a_chatml_style_template() {
let tmpl = "{%- for m in messages -%}<|im_start|>{{ m['role'] }}\n{{ m['content'] }}<|im_end|>\n{% endfor %}{%- if add_generation_prompt -%}<|im_start|>assistant\n{% endif %}";
let out = render_chat_template(tmpl, &msgs(), &vars(false)).unwrap();
assert_eq!(
out,
"<|im_start|>system\nBe brief.<|im_end|>\n<|im_start|>user\nHi<|im_end|>\n<|im_start|>assistant\n"
);
}
#[test]
fn enable_thinking_reaches_the_template() {
let tmpl = "{%- set enable_thinking = enable_thinking | default(false) -%}\
{%- if enable_thinking -%}<|think|>{%- endif -%}done";
let on = render_chat_template(tmpl, &msgs(), &vars(true)).unwrap();
let off = render_chat_template(tmpl, &msgs(), &vars(false)).unwrap();
assert_eq!(on, "<|think|>done");
assert_eq!(off, "done");
}
#[test]
fn tools_are_none_so_templates_do_not_emit_a_native_tool_dialect() {
let tmpl = "{%- if tools -%}native{%- else -%}absent{%- endif -%}";
let out = render_chat_template(tmpl, &msgs(), &vars(false)).unwrap();
assert_eq!(out, "absent");
}
#[test]
fn get_method_returns_default_for_missing_key() {
let tmpl = "{{ messages[0].get('role') }}|{{ messages[0].get('nope') }}|{{ messages[0].get('nope', 'fb') }}";
let out = render_chat_template(tmpl, &msgs(), &vars(false)).unwrap();
assert_eq!(out, "system|none|fb");
}
#[test]
fn split_method_splits_on_a_separator() {
let tmpl = "{%- for p in 'a<channel|>b'.split('<channel|>') -%}[{{ p }}]{%- endfor -%}";
let out = render_chat_template(tmpl, &msgs(), &vars(false)).unwrap();
assert_eq!(out, "[a][b]");
}
#[test]
fn supports_macros_namespaces_and_dictsort() {
let tmpl = "{%- macro emit(k, v) -%}{{ k }}={{ v }};{%- endmacro -%}\
{%- set ns = namespace(n=0) -%}\
{%- for k, v in {'b': 2, 'a': 1} | dictsort -%}\
{%- set ns.n = ns.n + 1 -%}{{ emit(k, v) }}\
{%- endfor -%}count={{ ns.n }}";
let out = render_chat_template(tmpl, &msgs(), &vars(false)).unwrap();
assert_eq!(out, "a=1;b=2;count=2");
}
#[test]
fn reports_an_error_for_a_broken_template() {
assert!(render_chat_template("{% for %}", &msgs(), &vars(false)).is_err());
}
}