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> {
fn text<'a>(value: &'a JValue, method: &str) -> Result<&'a str, Error> {
value.as_str().ok_or_else(|| {
Error::new(
ErrorKind::InvalidOperation,
format!("{method}() called on a non-string"),
)
})
}
fn arg(args: &[JValue]) -> Option<&str> {
args.first().and_then(|a| a.as_str())
}
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 = text(value, method)?;
let parts: Vec<JValue> = match arg(args) {
Some(sep) => text.split(sep).map(JValue::from).collect(),
None => text.split_whitespace().map(JValue::from).collect(),
};
Ok(JValue::from(parts))
}
"startswith" => Ok(JValue::from(arg(args).is_some_and(|prefix| {
text(value, method).is_ok_and(|t| t.starts_with(prefix))
}))),
"endswith" => Ok(JValue::from(arg(args).is_some_and(|suffix| {
text(value, method).is_ok_and(|t| t.ends_with(suffix))
}))),
"strip" | "lstrip" | "rstrip" => {
let text = text(value, method)?;
let cut: &dyn Fn(&str) -> &str = &|s: &str| match arg(args) {
Some(chars) => match method {
"lstrip" => s.trim_start_matches(|c| chars.contains(c)),
"rstrip" => s.trim_end_matches(|c| chars.contains(c)),
_ => s.trim_matches(|c| chars.contains(c)),
},
None => match method {
"lstrip" => s.trim_start(),
"rstrip" => s.trim_end(),
_ => s.trim(),
},
};
Ok(JValue::from(cut(text)))
}
_ => Err(Error::from(ErrorKind::UnknownMethod)),
}
}
pub(crate) struct TemplateVars<'a> {
pub tools: Option<&'a Value>,
pub enable_thinking: Option<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 => match vars.tools {
Some(tools) => JValue::from_serialize(tools),
None => JValue::from(()),
},
enable_thinking => match vars.enable_thinking {
Some(on) => JValue::from(on),
None => JValue::UNDEFINED,
},
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: Option<bool>) -> TemplateVars<'a> {
TemplateVars {
tools: None,
enable_thinking: thinking,
bos_token: "",
}
}
fn vars_with_tools<'a>(tools: &'a Value) -> TemplateVars<'a> {
TemplateVars {
tools: Some(tools),
enable_thinking: None,
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(None)).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(Some(true))).unwrap();
let off = render_chat_template(tmpl, &msgs(), &vars(Some(false))).unwrap();
assert_eq!(on, "<|think|>done");
assert_eq!(off, "done");
}
#[test]
fn an_unspecified_thinking_preference_leaves_the_template_default_alone() {
let tmpl = "{%- if enable_thinking is defined and enable_thinking is false -%}\
off{%- else -%}on{%- endif -%}";
assert_eq!(
render_chat_template(tmpl, &msgs(), &vars(None)).unwrap(),
"on"
);
assert_eq!(
render_chat_template(tmpl, &msgs(), &vars(Some(false))).unwrap(),
"off"
);
}
#[test]
fn tools_are_absent_unless_the_caller_passes_them() {
let tmpl = "{%- if tools -%}native{%- else -%}absent{%- endif -%}";
let out = render_chat_template(tmpl, &msgs(), &vars(None)).unwrap();
assert_eq!(out, "absent");
}
#[test]
fn a_tool_turn_hands_the_template_its_native_tools() {
let tools = json!([{"type": "function", "function": {"name": "list"}}]);
let tmpl = "{%- for t in tools -%}{{ t.function.name }}{%- endfor -%}";
let out = render_chat_template(tmpl, &msgs(), &vars_with_tools(&tools)).unwrap();
assert_eq!(out, "list");
}
#[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(None)).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(None)).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(None)).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(None)).is_err());
}
}