use std::path::Path;
use minijinja::value::{Kwargs, Value as JValue, ValueKind};
use minijinja::{Environment, Error as JError, ErrorKind as JErrorKind, State};
use serde::Serialize;
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)]
pub struct TemplateCapabilities {
pub supports_tools: bool,
pub supports_tool_calls: bool,
pub supports_system_role: bool,
pub supports_parallel_tool_calls: bool,
pub supports_tool_call_id: bool,
pub requires_typed_content: bool,
pub supports_single_turn: bool,
}
#[derive(Debug, Clone)]
pub struct TemplateError {
pub message: String,
pub stage: TemplateStage,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TemplateStage {
Compile,
Render,
}
impl std::fmt::Display for TemplateError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let what = match self.stage {
TemplateStage::Compile => "chat template compile",
TemplateStage::Render => "chat template render",
};
write!(f, "{what}: {}", self.message)
}
}
impl std::error::Error for TemplateError {}
impl TemplateError {
fn from_jinja(stage: TemplateStage, e: &JError) -> Self {
let mut message = e.to_string();
let mut src: Option<&(dyn std::error::Error + 'static)> = std::error::Error::source(e);
while let Some(e) = src {
message.push_str(": ");
message.push_str(&e.to_string());
src = std::error::Error::source(e);
}
Self { message, stage }
}
}
#[derive(Debug, Clone, Default)]
pub struct RenderInputs {
pub messages: serde_json::Value,
pub tools: Option<serde_json::Value>,
pub documents: Option<serde_json::Value>,
pub add_generation_prompt: bool,
pub now: Option<i64>,
pub extra: serde_json::Map<String, serde_json::Value>,
}
#[derive(Clone)]
pub struct ChatTemplate {
env: Environment<'static>,
source: String,
caps: TemplateCapabilities,
}
const TEMPLATE_NAME: &str = "chat";
impl ChatTemplate {
pub fn new(source: impl Into<String>) -> Result<Self, TemplateError> {
let source = source.into();
let mut env = create_env();
env.add_template_owned(TEMPLATE_NAME, rewrite_generation_blocks(&source))
.map_err(|e| TemplateError::from_jinja(TemplateStage::Compile, &e))?;
let caps = detect_capabilities(&source);
Ok(Self { env, source, caps })
}
pub fn source(&self) -> &str {
&self.source
}
pub fn capabilities(&self) -> TemplateCapabilities {
self.caps
}
pub fn render(&self, inputs: &RenderInputs) -> Result<String, TemplateError> {
let tmpl = self
.env
.get_template(TEMPLATE_NAME)
.map_err(|e| TemplateError::from_jinja(TemplateStage::Compile, &e))?;
let mut ctx = serde_json::Map::new();
ctx.insert(
"messages".into(),
match &inputs.messages {
serde_json::Value::Null => serde_json::Value::Array(Vec::new()),
other => other.clone(),
},
);
ctx.insert(
"tools".into(),
inputs.tools.clone().unwrap_or(serde_json::Value::Null),
);
ctx.insert(
"documents".into(),
inputs.documents.clone().unwrap_or(serde_json::Value::Null),
);
ctx.insert(
"add_generation_prompt".into(),
serde_json::Value::Bool(inputs.add_generation_prompt),
);
if let Some(now) = inputs.now {
ctx.insert("now".into(), serde_json::Value::from(now));
}
for (k, v) in &inputs.extra {
ctx.insert(k.clone(), v.clone());
}
tmpl.render(JValue::from_serialize(serde_json::Value::Object(ctx)))
.map_err(|e| TemplateError::from_jinja(TemplateStage::Render, &e))
}
}
pub struct CheckpointChat {
pub template: ChatTemplate,
pub special_tokens: serde_json::Map<String, serde_json::Value>,
}
const SPECIAL_TOKEN_KEYS: &[&str] = &[
"bos_token",
"eos_token",
"unk_token",
"sep_token",
"pad_token",
"cls_token",
"mask_token",
"additional_special_tokens",
];
pub fn load_checkpoint_chat(dir: &Path) -> Result<Option<CheckpointChat>, TemplateError> {
let tok_cfg: Option<serde_json::Value> =
std::fs::read_to_string(dir.join("tokenizer_config.json"))
.ok()
.and_then(|s| serde_json::from_str(&s).ok());
let source = match std::fs::read_to_string(dir.join("chat_template.jinja")) {
Ok(s) => Some(s),
Err(_) => tok_cfg
.as_ref()
.and_then(|c| c.get("chat_template"))
.and_then(template_from_config_value),
};
let Some(source) = source else {
return Ok(None);
};
let mut special_tokens = serde_json::Map::new();
for file in ["tokenizer_config.json", "special_tokens_map.json"] {
let Ok(raw) = std::fs::read_to_string(dir.join(file)) else {
continue;
};
let Ok(v) = serde_json::from_str::<serde_json::Value>(&raw) else {
continue;
};
for key in SPECIAL_TOKEN_KEYS {
if let Some(tok) = v.get(*key)
&& let Some(normalized) = normalize_special_token(tok)
{
special_tokens.insert((*key).to_string(), normalized);
}
}
}
Ok(Some(CheckpointChat {
template: ChatTemplate::new(source)?,
special_tokens,
}))
}
fn template_from_config_value(v: &serde_json::Value) -> Option<String> {
match v {
serde_json::Value::String(s) => Some(s.clone()),
serde_json::Value::Array(entries) => entries
.iter()
.find(|e| e.get("name").and_then(|n| n.as_str()) == Some("default"))
.or_else(|| entries.first())
.and_then(|e| e.get("template"))
.and_then(|t| t.as_str())
.map(str::to_string),
_ => None,
}
}
fn normalize_special_token(v: &serde_json::Value) -> Option<serde_json::Value> {
match v {
serde_json::Value::String(_) => Some(v.clone()),
serde_json::Value::Object(o) => o.get("content").filter(|c| c.is_string()).cloned(),
serde_json::Value::Array(items) => Some(serde_json::Value::Array(
items.iter().filter_map(normalize_special_token).collect(),
)),
_ => None,
}
}
fn detect_capabilities(source: &str) -> TemplateCapabilities {
let mut caps = TemplateCapabilities::default();
let env = create_env();
let compilable = rewrite_generation_blocks(source);
if let Ok(tmpl) = env.template_from_str(&compilable) {
if tmpl.undeclared_variables(true).contains("tools") {
caps.supports_tools = true;
}
const PROBE: &str = "test content";
let try_render = |content: serde_json::Value| -> bool {
let ctx = serde_json::json!({
"messages": [{"role": "user", "content": content}],
"add_generation_prompt": false,
"tools": [],
"documents": null,
});
tmpl.render(JValue::from_serialize(&ctx))
.map(|s| s.contains(PROBE))
.unwrap_or(false)
};
let string_works = try_render(serde_json::json!(PROBE));
let typed_works = try_render(serde_json::json!([{"type": "text", "text": PROBE}]));
if !string_works && typed_works {
caps.requires_typed_content = true;
}
}
if source.contains("tool_calls") {
caps.supports_tool_calls = true;
}
if source.contains("tool_call_id") {
caps.supports_tool_call_id = true;
}
if !caps.supports_tools && source.contains("tools") {
caps.supports_tools = true;
}
if source.contains("system") {
caps.supports_system_role = true;
}
if caps.supports_tool_calls && source.contains("for") {
caps.supports_parallel_tool_calls = true;
}
if source.contains("is_appending_to_prefill") {
caps.supports_single_turn = true;
}
caps
}
fn python_str(v: &JValue) -> String {
if v.is_undefined() {
return String::new(); }
if v.is_none() {
return "None".into();
}
match v.kind() {
ValueKind::String => v.as_str().unwrap_or_default().to_owned(),
ValueKind::Bool => bool_repr(v).into(),
ValueKind::Seq | ValueKind::Iterable | ValueKind::Map => python_repr(v),
_ => v.to_string(),
}
}
fn python_repr(v: &JValue) -> String {
if v.is_none() {
return "None".into();
}
match v.kind() {
ValueKind::String => python_quote(v.as_str().unwrap_or_default()),
ValueKind::Bool => bool_repr(v).into(),
ValueKind::Seq | ValueKind::Iterable => {
let items = v
.try_iter()
.map(|it| it.map(|x| python_repr(&x)).collect::<Vec<_>>())
.unwrap_or_default();
format!("[{}]", items.join(", "))
}
ValueKind::Map => {
let entries = v
.try_iter()
.map(|it| {
it.map(|k| {
let val = v.get_item(&k).unwrap_or(JValue::UNDEFINED);
format!("{}: {}", python_repr(&k), python_repr(&val))
})
.collect::<Vec<_>>()
})
.unwrap_or_default();
format!("{{{}}}", entries.join(", "))
}
_ => v.to_string(),
}
}
fn bool_repr(v: &JValue) -> &'static str {
if v.is_true() { "True" } else { "False" }
}
fn python_quote(s: &str) -> String {
let quote = if s.contains('\'') && !s.contains('"') {
'"'
} else {
'\''
};
let mut out = String::with_capacity(s.len() + 2);
out.push(quote);
for c in s.chars() {
match c {
'\\' => out.push_str("\\\\"),
'\n' => out.push_str("\\n"),
'\r' => out.push_str("\\r"),
'\t' => out.push_str("\\t"),
c if c == quote => {
out.push('\\');
out.push(c);
}
c if (c as u32) < 0x20 || c as u32 == 0x7f => {
out.push_str(&format!("\\x{:02x}", c as u32));
}
c => out.push(c),
}
}
out.push(quote);
out
}
fn rewrite_generation_blocks(source: &str) -> String {
let mut out = String::with_capacity(source.len());
let mut i = 0;
while i < source.len() {
let Some(rel) = source[i..].find("{%") else {
out.push_str(&source[i..]);
return out;
};
let start = i + rel;
out.push_str(&source[i..start]);
let Some(rel_end) = source[start..].find("%}") else {
out.push_str(&source[start..]);
return out;
};
let end = start + rel_end + 2;
let inner = &source[start + 2..end - 2];
let (lead_dash, rest) = match inner.strip_prefix('-') {
Some(r) => (true, r),
None => (false, inner),
};
let (trail_dash, rest) = match rest.strip_suffix('-') {
Some(r) => (true, r),
None => (false, rest),
};
match rest.trim() {
"generation" => push_stmt(&mut out, lead_dash, "if true", trail_dash),
"endgeneration" => push_stmt(&mut out, lead_dash, "endif", trail_dash),
_ => out.push_str(&source[start..end]),
}
i = end;
}
out
}
fn push_stmt(out: &mut String, lead_dash: bool, stmt: &str, trail_dash: bool) {
out.push_str(if lead_dash { "{%- " } else { "{% " });
out.push_str(stmt);
out.push_str(if trail_dash { " -%}" } else { " %}" });
}
fn create_env() -> Environment<'static> {
let mut env = Environment::new();
env.set_trim_blocks(true);
env.set_lstrip_blocks(true);
env.set_unknown_method_callback(minijinja_contrib::pycompat::unknown_method_callback);
env.set_formatter(|out, state, value| {
match value.kind() {
ValueKind::Bool
| ValueKind::None
| ValueKind::Map
| ValueKind::Seq
| ValueKind::Iterable => {
minijinja::escape_formatter(out, state, &JValue::from(python_str(value)))
}
_ => minijinja::escape_formatter(out, state, value),
}
});
env.add_filter("string", |v: JValue| python_str(&v));
env.add_function("strftime_now", strftime_now);
env.add_function("raise_exception", raise_exception);
env.add_filter("tojson", tojson);
env.add_filter("lstrip", lstrip);
env.add_filter("rstrip", rstrip);
env
}
fn strftime_now(state: &State, format: String) -> Result<String, JError> {
let now = state.lookup("now");
let timestamp = match now {
Some(v) if !v.is_undefined() && !v.is_none() => i64::try_from(v)
.map_err(|_| JError::new(JErrorKind::InvalidOperation, "`now` must be unix seconds"))?,
_ => wall_clock_seconds()?,
};
let dt = chrono::DateTime::from_timestamp(timestamp, 0)
.ok_or_else(|| JError::new(JErrorKind::InvalidOperation, "timestamp out of range"))?;
Ok(dt.format(&format).to_string())
}
#[cfg(not(target_arch = "wasm32"))]
fn wall_clock_seconds() -> Result<i64, JError> {
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.map(|d| d.as_secs() as i64)
.map_err(|e| JError::new(JErrorKind::InvalidOperation, e.to_string()))
}
#[cfg(target_arch = "wasm32")]
fn wall_clock_seconds() -> Result<i64, JError> {
Err(JError::new(
JErrorKind::InvalidOperation,
"strftime_now needs an explicit `now` on this target (no wall clock)",
))
}
fn raise_exception(msg: String) -> Result<String, JError> {
Err(JError::new(JErrorKind::InvalidOperation, msg))
}
struct SpaceFormatter;
impl serde_json::ser::Formatter for SpaceFormatter {
fn begin_array_value<W>(&mut self, writer: &mut W, first: bool) -> std::io::Result<()>
where
W: ?Sized + std::io::Write,
{
if !first {
writer.write_all(b", ")?;
}
Ok(())
}
fn begin_object_key<W>(&mut self, writer: &mut W, first: bool) -> std::io::Result<()>
where
W: ?Sized + std::io::Write,
{
if !first {
writer.write_all(b", ")?;
}
Ok(())
}
fn begin_object_value<W>(&mut self, writer: &mut W) -> std::io::Result<()>
where
W: ?Sized + std::io::Write,
{
writer.write_all(b": ")
}
}
fn tojson(value: JValue, kwargs: Kwargs) -> Result<String, JError> {
let indent: Option<usize> = kwargs.get("indent").unwrap_or(None);
let ensure_ascii: bool = kwargs.get("ensure_ascii").unwrap_or(false);
kwargs.assert_all_used()?;
let mut buf = Vec::new();
let err = |e: String| JError::new(JErrorKind::InvalidOperation, e);
match indent {
Some(n) => {
let pad = " ".repeat(n);
let fmt = serde_json::ser::PrettyFormatter::with_indent(pad.as_bytes());
let mut ser = serde_json::Serializer::with_formatter(&mut buf, fmt);
value.serialize(&mut ser).map_err(|e| err(e.to_string()))?;
}
None => {
let mut ser = serde_json::Serializer::with_formatter(&mut buf, SpaceFormatter);
value.serialize(&mut ser).map_err(|e| err(e.to_string()))?;
}
}
let out = String::from_utf8(buf).map_err(|e| err(e.to_string()))?;
Ok(if ensure_ascii {
escape_non_ascii(&out)
} else {
out
})
}
fn escape_non_ascii(s: &str) -> String {
if s.is_ascii() {
return s.to_string();
}
let mut out = String::with_capacity(s.len());
for c in s.chars() {
if c.is_ascii() {
out.push(c);
} else {
let mut buf = [0u16; 2];
for unit in c.encode_utf16(&mut buf) {
out.push_str(&format!("\\u{unit:04x}"));
}
}
}
out
}
fn lstrip(s: std::borrow::Cow<'_, str>, chars: Option<std::borrow::Cow<'_, str>>) -> String {
match chars {
Some(chars) => {
let set = chars.chars().collect::<Vec<_>>();
s.trim_start_matches(&set[..]).to_string()
}
None => s.trim_start().to_string(),
}
}
fn rstrip(s: std::borrow::Cow<'_, str>, chars: Option<std::borrow::Cow<'_, str>>) -> String {
match chars {
Some(chars) => {
let set = chars.chars().collect::<Vec<_>>();
s.trim_end_matches(&set[..]).to_string()
}
None => s.trim_end().to_string(),
}
}
#[cfg(feature = "server")]
impl From<TemplateError> for super::error::ApiError {
fn from(e: TemplateError) -> Self {
super::error::ApiError::invalid_request(e.to_string()).with_param("messages")
}
}
#[cfg(feature = "server")]
pub fn wire_messages(
msgs: &[super::types::ChatMessage],
tools_preamble: Option<&str>,
caps: TemplateCapabilities,
) -> serde_json::Value {
use serde_json::{Map, Value};
fn content_value(text: &str, typed: bool) -> Value {
if typed {
let mut block = Map::new();
block.insert("type".into(), Value::from("text"));
block.insert("text".into(), Value::from(text));
Value::Array(vec![Value::Object(block)])
} else {
Value::from(text)
}
}
let typed = caps.requires_typed_content;
let mut out: Vec<Value> = Vec::with_capacity(msgs.len() + 1);
let first_is_system = msgs.first().is_some_and(|m| m.role == "system");
if let Some(block) = tools_preamble
&& !first_is_system
{
let mut sys = Map::new();
sys.insert("role".into(), Value::from("system"));
sys.insert("content".into(), content_value(block, typed));
out.push(Value::Object(sys));
}
for (i, m) in msgs.iter().enumerate() {
let mut obj = Map::new();
obj.insert("role".into(), Value::from(m.role.as_str()));
let mut content = m.content.clone();
if i == 0
&& first_is_system
&& let Some(block) = tools_preamble
{
if !content.is_empty() {
content.push_str("\n\n");
}
content.push_str(block);
}
obj.insert("content".into(), content_value(&content, typed));
if !m.tool_calls.is_empty() {
let calls = m
.tool_calls
.iter()
.map(|tc| {
let raw = tc.function.arguments.trim();
let args = if raw.is_empty() {
Value::Object(Map::new())
} else {
serde_json::from_str(raw).unwrap_or_else(|_| Value::from(raw))
};
let mut func = Map::new();
func.insert("name".into(), Value::from(tc.function.name.as_str()));
func.insert("arguments".into(), args);
let mut call = Map::new();
if let Some(id) = &tc.id {
call.insert("id".into(), Value::from(id.as_str()));
}
call.insert(
"type".into(),
Value::from(tc.kind.as_deref().unwrap_or("function")),
);
call.insert("function".into(), Value::Object(func));
Value::Object(call)
})
.collect::<Vec<_>>();
obj.insert("tool_calls".into(), Value::Array(calls));
}
if let Some(id) = &m.tool_call_id {
obj.insert("tool_call_id".into(), Value::from(id.as_str()));
}
if let Some(name) = &m.name {
obj.insert("name".into(), Value::from(name.as_str()));
}
out.push(Value::Object(obj));
}
Value::Array(out)
}
#[cfg(test)]
mod tests {
use super::*;
fn render(source: &str, inputs: &RenderInputs) -> Result<String, TemplateError> {
ChatTemplate::new(source).expect("compile").render(inputs)
}
#[test]
fn renders_a_basic_conversation() {
let out = render(
"Hello {{ messages[0].content }}",
&RenderInputs {
messages: serde_json::json!([{"role": "user", "content": "World"}]),
..Default::default()
},
)
.expect("render");
assert_eq!(out, "Hello World");
}
#[test]
fn tojson_spacing_matches_python_json_dumps() {
let mut extra = serde_json::Map::new();
extra.insert("data".into(), serde_json::json!({"a": 1, "b": [1, 2]}));
let out = render(
"{{ data | tojson }}",
&RenderInputs {
extra,
..Default::default()
},
)
.expect("render");
assert_eq!(out, r#"{"a": 1, "b": [1, 2]}"#);
}
#[test]
fn tojson_indent_matches_python_pretty_printing() {
let mut extra = serde_json::Map::new();
extra.insert("data".into(), serde_json::json!({"a": [1], "e": {}}));
let out = render(
"{{ data | tojson(indent=4) }}",
&RenderInputs {
extra,
..Default::default()
},
)
.expect("render");
assert_eq!(out, "{\n \"a\": [\n 1\n ],\n \"e\": {}\n}");
}
#[test]
fn tojson_ensure_ascii_escapes_like_cpython() {
let mut extra = serde_json::Map::new();
extra.insert("data".into(), serde_json::json!({"s": "é🙂"}));
let raw = render(
"{{ data | tojson }}",
&RenderInputs {
extra: extra.clone(),
..Default::default()
},
)
.expect("render");
assert_eq!(raw, "{\"s\": \"é🙂\"}");
let escaped = render(
"{{ data | tojson(ensure_ascii=True) }}",
&RenderInputs {
extra,
..Default::default()
},
)
.expect("render");
assert_eq!(escaped, r#"{"s": "\u00e9\ud83d\ude42"}"#);
}
#[test]
fn an_unsupported_tojson_keyword_is_loud() {
let mut extra = serde_json::Map::new();
extra.insert("data".into(), serde_json::json!([1]));
let err = render(
"{{ data | tojson(separators=',') }}",
&RenderInputs {
extra,
..Default::default()
},
)
.expect_err("unknown keyword must fail");
assert_eq!(err.stage, TemplateStage::Render);
}
#[test]
fn raise_exception_is_an_error_carrying_the_message() {
let err = render("{{ raise_exception('nope') }}", &RenderInputs::default())
.expect_err("must fail");
assert!(err.message.contains("nope"), "{err}");
assert_eq!(err.stage, TemplateStage::Render);
}
#[test]
fn strftime_now_formats_the_supplied_instant_in_utc() {
let out = render(
"{{ strftime_now('%Y-%m-%d') }}",
&RenderInputs {
now: Some(1735689600),
..Default::default()
},
)
.expect("render");
assert_eq!(out, "2025-01-01");
}
#[test]
fn lstrip_and_rstrip_take_a_char_set() {
let out = render(
"{{ ' bar '|lstrip }}|{{ ' baz '|rstrip }}|{{ '1212foo12'|lstrip('12') }}",
&RenderInputs::default(),
)
.expect("render");
assert_eq!(out, "bar | baz|foo12");
}
#[test]
fn python_string_methods_resolve() {
let out = render(
"{{ 'a,b'.split(',') | join('|') }}/{{ 'xy'.startswith('x') }}/{{ 'ab'.replace('a','c') }}",
&RenderInputs::default(),
)
.expect("render");
assert_eq!(out, "a|b/True/cb");
}
#[test]
fn values_stringify_the_way_cpython_does() {
let mut extra = serde_json::Map::new();
extra.insert(
"d".into(),
serde_json::json!({"type": "function", "ok": true, "n": null, "xs": [1, "a"]}),
);
extra.insert("s".into(), serde_json::json!("plain"));
extra.insert("q".into(), serde_json::json!(["it's", "say \"hi\""]));
let out = render(
"{{ d }}|{{ d | string }}|{{ s }}|{{ true }}|{{ none }}|{{ q }}",
&RenderInputs {
extra,
..Default::default()
},
)
.expect("render");
assert_eq!(
out,
"{'type': 'function', 'ok': True, 'n': None, 'xs': [1, 'a']}\
|{'type': 'function', 'ok': True, 'n': None, 'xs': [1, 'a']}\
|plain|True|None|[\"it's\", 'say \"hi\"']"
);
}
#[test]
fn the_none_test_follows_jinja2_semantics() {
let out = render(
"{{ tools is none }}/{{ nothing is none }}/{{ nothing is undefined }}",
&RenderInputs::default(),
)
.expect("render");
assert_eq!(out, "True/False/True");
}
#[test]
fn a_broken_template_fails_to_compile_without_panicking() {
let err = ChatTemplate::new("{% for x in %}")
.err()
.expect("must fail");
assert_eq!(err.stage, TemplateStage::Compile);
}
}