use std::{
collections::{BTreeMap, BTreeSet},
fs::read_to_string,
ops::{Deref, DerefMut},
path::Path,
str::FromStr,
};
use minijinja::{
value::{Kwargs, Value},
Environment, Template,
};
use serde::Serialize;
use sha2::{Digest, Sha256};
use tokenizers::Encoding;
use crate::error::Error;
const DEFAULT_CHAT_TEMPLATE_NAME: &str = "default";
const TOOL_USE_CHAT_TEMPLATE_NAME: &str = "tool_use";
pub fn vocabulary_fingerprint(tokenizer: &tokenizers::Tokenizer) -> [u8; 32] {
let vocabulary_size = tokenizer.get_vocab_size(true);
let mut hasher = Sha256::new();
hasher.update(b"eredu-token-id-vocabulary-v1");
hasher.update((vocabulary_size as u64).to_le_bytes());
for token_id in 0..vocabulary_size {
hasher.update((token_id as u64).to_le_bytes());
match tokenizer.id_to_token(token_id as u32) {
Some(token) => {
hasher.update((token.len() as u64).to_le_bytes());
hasher.update(token.as_bytes());
}
None => hasher.update(u64::MAX.to_le_bytes()),
}
}
hasher.finalize().into()
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ModelChatTemplate {
Single(String),
Named(BTreeMap<String, String>),
}
impl From<String> for ModelChatTemplate {
fn from(template: String) -> Self {
Self::Single(template)
}
}
impl From<&str> for ModelChatTemplate {
fn from(template: &str) -> Self {
Self::Single(template.to_owned())
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub enum ChatTemplateIdentity {
Single,
Named(String),
}
#[derive(Debug, Clone)]
pub struct SelectedChatTemplate<'a> {
template: &'a str,
identity: ChatTemplateIdentity,
}
impl SelectedChatTemplate<'_> {
pub fn template(&self) -> &str {
self.template
}
pub fn identity(&self) -> &ChatTemplateIdentity {
&self.identity
}
}
impl ModelChatTemplate {
pub fn select(
&self,
tools: Option<&[serde_json::Value]>,
) -> Result<SelectedChatTemplate<'_>, Error> {
match self {
Self::Single(template) => Ok(SelectedChatTemplate {
template,
identity: ChatTemplateIdentity::Single,
}),
Self::Named(templates) => {
let selected_name = if tools.is_some_and(|tools| !tools.is_empty())
&& templates.contains_key(TOOL_USE_CHAT_TEMPLATE_NAME)
{
TOOL_USE_CHAT_TEMPLATE_NAME
} else {
DEFAULT_CHAT_TEMPLATE_NAME
};
let template =
templates
.get(selected_name)
.ok_or_else(|| Error::AmbiguousChatTemplate {
available: templates.keys().cloned().collect(),
})?;
Ok(SelectedChatTemplate {
template,
identity: ChatTemplateIdentity::Named(selected_name.to_owned()),
})
}
}
}
}
pub struct Tokenizer {
inner: tokenizers::Tokenizer,
env: Environment<'static>,
template_kwargs: serde_json::Map<String, serde_json::Value>,
}
struct TemplateKwargs<'defaults, 'overrides> {
defaults: Option<&'defaults serde_json::Map<String, serde_json::Value>>,
overrides: Option<&'overrides serde_json::Map<String, serde_json::Value>>,
}
impl FromStr for Tokenizer {
type Err = tokenizers::Error;
fn from_str(s: &str) -> Result<Self, Self::Err> {
tokenizers::Tokenizer::from_str(s).map(Self::from_tokenizer)
}
}
impl Tokenizer {
pub fn from_tokenizer(tokenizer: tokenizers::Tokenizer) -> Self {
let mut env = Environment::new();
env.set_unknown_method_callback(minijinja_contrib::pycompat::unknown_method_callback);
env.add_filter("tojson", hugging_face_tojson);
Self {
inner: tokenizer,
env,
template_kwargs: serde_json::Map::new(),
}
}
pub fn set_template_kwargs(
&mut self,
template_kwargs: serde_json::Map<String, serde_json::Value>,
) {
self.template_kwargs = template_kwargs;
}
pub fn template_kwargs(&self) -> &serde_json::Map<String, serde_json::Value> {
&self.template_kwargs
}
pub fn from_file(file: impl AsRef<Path>) -> tokenizers::Result<Self> {
tokenizers::Tokenizer::from_file(file).map(Self::from_tokenizer)
}
pub fn from_bytes(bytes: impl AsRef<[u8]>) -> tokenizers::Result<Self> {
tokenizers::Tokenizer::from_bytes(bytes).map(Self::from_tokenizer)
}
pub fn apply_chat_template<'a, I, R, T>(
&'a mut self,
model_template: impl Into<ModelChatTemplate>,
args: ApplyChatTemplateArgs<'a, I, R, T>,
) -> Result<Vec<String>, Error>
where
I: IntoIterator<Item = Chat<'a, R, T>>,
R: Serialize + 'a,
T: Serialize + 'a,
{
apply_chat_template_with_default_kwargs(
&mut self.env,
model_template.into(),
args,
Some(&self.template_kwargs),
)
}
pub fn apply_chat_template_and_encode<'a, I, R, T>(
&mut self,
model_template: impl Into<ModelChatTemplate>,
args: ApplyChatTemplateArgs<'a, I, R, T>,
) -> Result<Vec<Encoding>, Error>
where
I: IntoIterator<Item = Chat<'a, R, T>>,
R: Serialize + 'a,
T: Serialize + 'a,
{
let Self {
inner,
env,
template_kwargs,
} = self;
let rendered_chats = apply_chat_template_with_default_kwargs(
env,
model_template.into(),
args,
Some(template_kwargs),
)?;
inner
.encode_batch(rendered_chats, false)
.map_err(Into::into)
}
pub fn apply_chat_template_json<'a, I>(
&mut self,
model_template: impl Into<ModelChatTemplate>,
conversations: I,
tools: Option<&'a [serde_json::Value]>,
model_id: &'a str,
add_generation_prompt: bool,
template_kwargs: Option<&'a serde_json::Map<String, serde_json::Value>>,
) -> Result<Vec<String>, Error>
where
I: IntoIterator<Item = Vec<serde_json::Value>>,
{
apply_chat_template_json_with_default_kwargs(
&mut self.env,
model_template.into(),
conversations,
tools,
model_id,
add_generation_prompt,
TemplateKwargs {
defaults: Some(&self.template_kwargs),
overrides: template_kwargs,
},
)
}
}
fn hugging_face_tojson(
value: &Value,
indent: Option<Value>,
kwargs: Kwargs,
) -> Result<Value, minijinja::Error> {
let _: Option<Value> = kwargs.get("separators")?;
let _: Option<bool> = kwargs.get("sort_keys")?;
minijinja::filters::tojson(value, indent, kwargs)
}
impl Deref for Tokenizer {
type Target = tokenizers::Tokenizer;
fn deref(&self) -> &Self::Target {
&self.inner
}
}
impl DerefMut for Tokenizer {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.inner
}
}
#[derive(Debug, Clone, Copy, Serialize)]
#[serde(rename_all = "lowercase")]
pub enum Role {
System,
User,
Assistant,
}
#[derive(Debug, Clone, Serialize)]
pub struct Conversation<R, T> {
pub role: R,
pub content: T,
}
#[derive(Debug, Clone, Serialize)]
#[serde(untagged)]
pub enum Chat<'a, R, T> {
Borrowed(&'a [Conversation<R, T>]),
Owned(Vec<Conversation<R, T>>),
}
impl<R, T> Deref for Chat<'_, R, T> {
type Target = [Conversation<R, T>];
fn deref(&self) -> &Self::Target {
match self {
Chat::Borrowed(conversations) => conversations,
Chat::Owned(conversations) => conversations,
}
}
}
impl<R, T> From<Vec<Conversation<R, T>>> for Chat<'_, R, T> {
fn from(value: Vec<Conversation<R, T>>) -> Self {
Chat::Owned(value)
}
}
impl<'a, R, T> From<&'a [Conversation<R, T>]> for Chat<'a, R, T> {
fn from(value: &'a [Conversation<R, T>]) -> Self {
Chat::Borrowed(value)
}
}
#[derive(Debug, Clone, Serialize)]
pub struct Document {
pub title: String,
pub text: String,
}
#[derive(Debug, Clone, Serialize)]
#[serde(transparent)]
pub struct JsonConversation(pub Vec<serde_json::Value>);
#[derive(Default)]
pub struct ApplyChatTemplateArgs<'a, I, R = Role, T = String>
where
I: IntoIterator<Item = Chat<'a, R, T>>,
R: Serialize + 'a,
T: Serialize + 'a,
{
pub conversations: I,
pub tools: Option<&'a [serde_json::Value]>,
pub documents: Option<&'a [Document]>,
pub model_id: &'a str,
pub chat_template_id: Option<&'a str>,
pub add_generation_prompt: Option<bool>,
pub continue_final_message: Option<bool>,
pub template_kwargs: Option<&'a serde_json::Map<String, serde_json::Value>>,
}
pub fn load_model_chat_template_from_str(
content: &str,
) -> std::io::Result<Option<ModelChatTemplate>> {
let config =
serde_json::from_str::<serde_json::Value>(content).map_err(std::io::Error::from)?;
let Some(value) = config.get("chat_template") else {
return Ok(None);
};
if value.is_null() {
return Ok(None);
}
if let Some(template) = value.as_str() {
return Ok(Some(ModelChatTemplate::Single(template.to_owned())));
}
let Some(entries) = value.as_array() else {
return Err(invalid_chat_template(
"expected a string or an array of named template entries".into(),
));
};
if entries.is_empty() {
return Err(invalid_chat_template(
"named template collection must not be empty".into(),
));
}
let mut templates = BTreeMap::new();
for (index, entry) in entries.iter().enumerate() {
let Some(entry) = entry.as_object() else {
return Err(invalid_chat_template(format!(
"entry {index} must be an object with string fields \"name\" and \"template\""
)));
};
if entry.len() != 2 || !entry.contains_key("name") || !entry.contains_key("template") {
return Err(invalid_chat_template(format!(
"entry {index} must contain exactly the fields \"name\" and \"template\""
)));
}
let Some(name) = entry.get("name").and_then(serde_json::Value::as_str) else {
return Err(invalid_chat_template(format!(
"entry {index} field \"name\" must be a string"
)));
};
if name.is_empty() {
return Err(invalid_chat_template(format!(
"entry {index} field \"name\" must not be empty"
)));
}
let Some(template) = entry.get("template").and_then(serde_json::Value::as_str) else {
return Err(invalid_chat_template(format!(
"entry {index} field \"template\" must be a string"
)));
};
if templates
.insert(name.to_owned(), template.to_owned())
.is_some()
{
return Err(invalid_chat_template(format!(
"duplicate named template {name:?}"
)));
}
}
Ok(Some(ModelChatTemplate::Named(templates)))
}
pub fn load_model_chat_template_from_file(
file: impl AsRef<Path>,
) -> std::io::Result<Option<ModelChatTemplate>> {
let content = read_to_string(file)?;
load_model_chat_template_from_str(&content)
}
fn invalid_chat_template(message: String) -> std::io::Error {
std::io::Error::new(
std::io::ErrorKind::InvalidData,
Error::InvalidChatTemplate(message),
)
}
pub fn chat_template_kwargs(
model_template: &str,
model_id: &str,
) -> Result<BTreeSet<String>, Error> {
let mut env = Environment::new();
env.set_unknown_method_callback(minijinja_contrib::pycompat::unknown_method_callback);
let compatible_template = normalize_chat_template(model_template);
env.add_template_owned(model_id.to_owned(), compatible_template)?;
let template = env.get_template(model_id)?;
let globals = env
.globals()
.map(|(name, _)| name.to_string())
.collect::<BTreeSet<_>>();
Ok(template
.undeclared_variables(false)
.into_iter()
.filter(|name| !STANDARD_CHAT_TEMPLATE_VARIABLES.contains(&name.as_str()))
.filter(|name| !globals.contains(name))
.collect())
}
const STANDARD_CHAT_TEMPLATE_VARIABLES: &[&str] = &[
"messages",
"tools",
"documents",
"add_generation_prompt",
"raise_exception",
"strftime_now",
];
pub fn apply_chat_template_json<'a, I>(
env: &mut Environment<'static>,
model_template: impl Into<ModelChatTemplate>,
conversations: I,
tools: Option<&'a [serde_json::Value]>,
model_id: &'a str,
add_generation_prompt: bool,
template_kwargs: Option<&'a serde_json::Map<String, serde_json::Value>>,
) -> Result<Vec<String>, Error>
where
I: IntoIterator<Item = Vec<serde_json::Value>>,
{
apply_chat_template_json_with_default_kwargs(
env,
model_template.into(),
conversations,
tools,
model_id,
add_generation_prompt,
TemplateKwargs {
defaults: None,
overrides: template_kwargs,
},
)
}
fn apply_chat_template_json_with_default_kwargs<'a, 'defaults, I>(
env: &mut Environment<'static>,
model_template: ModelChatTemplate,
conversations: I,
tools: Option<&'a [serde_json::Value]>,
model_id: &'a str,
add_generation_prompt: bool,
template_kwargs: TemplateKwargs<'defaults, 'a>,
) -> Result<Vec<String>, Error>
where
I: IntoIterator<Item = Vec<serde_json::Value>>,
{
let conversations = conversations.into_iter().map(|conversation| {
Chat::Owned(vec![Conversation {
role: serde_json::Value::Null,
content: JsonConversation(conversation),
}])
});
apply_chat_template_with_default_kwargs(
env,
model_template,
ApplyChatTemplateArgs {
conversations,
tools,
documents: None,
model_id,
chat_template_id: None,
add_generation_prompt: Some(add_generation_prompt),
continue_final_message: None,
template_kwargs: template_kwargs.overrides,
},
template_kwargs.defaults,
)
}
pub fn apply_chat_template<'a, I, R, T>(
env: &mut Environment<'static>,
model_template: impl Into<ModelChatTemplate>,
args: ApplyChatTemplateArgs<'a, I, R, T>,
) -> Result<Vec<String>, Error>
where
I: IntoIterator<Item = Chat<'a, R, T>>,
R: Serialize + 'a,
T: Serialize + 'a,
{
apply_chat_template_with_default_kwargs(env, model_template.into(), args, None)
}
fn apply_chat_template_with_default_kwargs<'a, I, R, T>(
env: &mut Environment<'static>,
model_template: ModelChatTemplate,
args: ApplyChatTemplateArgs<'a, I, R, T>,
default_template_kwargs: Option<&serde_json::Map<String, serde_json::Value>>,
) -> Result<Vec<String>, Error>
where
I: IntoIterator<Item = Chat<'a, R, T>>,
R: Serialize + 'a,
T: Serialize + 'a,
{
env.add_function("strftime_now", |format: &str| {
chrono::Local::now().format(format).to_string()
});
let ApplyChatTemplateArgs {
conversations,
tools,
documents,
model_id,
chat_template_id,
add_generation_prompt,
continue_final_message,
template_kwargs,
} = args;
let add_generation_prompt = add_generation_prompt.unwrap_or(false);
let continue_final_message = continue_final_message.unwrap_or(false);
let selected = model_template.select(tools)?;
let template = match chat_template_id {
Some(chat_template_id) => env.get_template(chat_template_id)?,
None => {
let selected_template_id = match selected.identity() {
ChatTemplateIdentity::Single => model_id.to_owned(),
ChatTemplateIdentity::Named(name) => {
format!("{model_id}::chat_template::{name}")
}
};
match env.get_template(&selected_template_id) {
Ok(template) => template,
Err(_) => {
let compatible_template = normalize_chat_template(selected.template());
env.add_template_owned(selected_template_id.clone(), compatible_template)?;
env.get_template(&selected_template_id)
.expect("Newly added template must be present")
}
}
}
};
render_jinja_template(
template,
conversations,
tools,
documents,
Some(add_generation_prompt),
Some(continue_final_message),
TemplateKwargs {
defaults: default_template_kwargs,
overrides: template_kwargs,
},
)
}
fn normalize_generation_blocks(template: &str) -> String {
let mut output = String::with_capacity(template.len());
let mut remaining = template;
while let Some(start) = remaining.find("{%") {
output.push_str(&remaining[..start]);
let statement = &remaining[start..];
let Some(end) = statement.find("%}") else {
output.push_str(statement);
return output;
};
let end = end + 2;
let tag = &statement[..end];
let body = tag[2..tag.len() - 2].trim().trim_matches('-').trim();
match body {
"generation" => output.push_str(&tag.replacen("generation", "if true", 1)),
"endgeneration" => output.push_str(&tag.replacen("endgeneration", "endif", 1)),
_ => output.push_str(tag),
}
remaining = &statement[end..];
}
output.push_str(remaining);
output
}
fn normalize_conditional_keyword_arguments(template: &str) -> String {
fn is_identifier(byte: u8) -> bool {
byte.is_ascii_alphanumeric() || byte == b'_'
}
fn keyword_at(bytes: &[u8], index: usize, keyword: &[u8]) -> bool {
bytes.get(index..index + keyword.len()) == Some(keyword)
&& (index == 0 || !is_identifier(bytes[index - 1]))
&& bytes
.get(index + keyword.len())
.is_none_or(|byte| !is_identifier(*byte))
}
fn is_fully_parenthesized(bytes: &[u8], start: usize, end: usize) -> bool {
let Some(first) = (start..end).find(|index| !bytes[*index].is_ascii_whitespace()) else {
return false;
};
let Some(last) = (start..end)
.rev()
.find(|index| !bytes[*index].is_ascii_whitespace())
else {
return false;
};
if bytes[first] != b'(' || bytes[last] != b')' {
return false;
}
let mut depth = 0usize;
let mut quote = None;
let mut escaped = false;
for (offset, byte) in bytes[first..=last].iter().copied().enumerate() {
if let Some(active_quote) = quote {
if escaped {
escaped = false;
} else if byte == b'\\' {
escaped = true;
} else if byte == active_quote {
quote = None;
}
continue;
}
match byte {
b'\'' | b'"' => quote = Some(byte),
b'(' => depth += 1,
b')' => {
depth = depth.saturating_sub(1);
if depth == 0 && first + offset != last {
return false;
}
}
_ => {}
}
}
depth == 0 && quote.is_none()
}
fn normalize_tag(tag: &str) -> String {
let bytes = tag.as_bytes();
let mut paren_depth = 0usize;
let mut bracket_depth = 0usize;
let mut brace_depth = 0usize;
let mut quote = None;
let mut escaped = false;
let mut insertions = Vec::new();
for (index, byte) in bytes.iter().copied().enumerate() {
if let Some(active_quote) = quote {
if escaped {
escaped = false;
} else if byte == b'\\' {
escaped = true;
} else if byte == active_quote {
quote = None;
}
continue;
}
match byte {
b'\'' | b'"' => quote = Some(byte),
b'(' => paren_depth += 1,
b')' => paren_depth = paren_depth.saturating_sub(1),
b'[' => bracket_depth += 1,
b']' => bracket_depth = bracket_depth.saturating_sub(1),
b'{' => brace_depth += 1,
b'}' => brace_depth = brace_depth.saturating_sub(1),
b'=' if paren_depth > 0
&& bytes.get(index.wrapping_sub(1)).is_some_and(|byte| {
*byte != b'=' && *byte != b'!' && *byte != b'<' && *byte != b'>'
})
&& bytes.get(index + 1) != Some(&b'=') =>
{
let Some(lhs_end) = (0..index)
.rev()
.find(|position| !bytes[*position].is_ascii_whitespace())
else {
continue;
};
if !is_identifier(bytes[lhs_end]) {
continue;
}
let base_paren = paren_depth;
let base_bracket = bracket_depth;
let base_brace = brace_depth;
let mut scan_paren = paren_depth;
let mut scan_bracket = bracket_depth;
let mut scan_brace = brace_depth;
let mut scan_quote = None;
let mut scan_escaped = false;
let mut found_if = false;
let mut found_else = false;
let mut end = bytes.len();
let mut cursor = index + 1;
while cursor < bytes.len() {
let current = bytes[cursor];
if let Some(active_quote) = scan_quote {
if scan_escaped {
scan_escaped = false;
} else if current == b'\\' {
scan_escaped = true;
} else if current == active_quote {
scan_quote = None;
}
cursor += 1;
continue;
}
match current {
b'\'' | b'"' => scan_quote = Some(current),
b'(' => scan_paren += 1,
b')' if scan_paren == base_paren
&& scan_bracket == base_bracket
&& scan_brace == base_brace =>
{
end = cursor;
break;
}
b')' => scan_paren = scan_paren.saturating_sub(1),
b'[' => scan_bracket += 1,
b']' => scan_bracket = scan_bracket.saturating_sub(1),
b'{' => scan_brace += 1,
b'}' => scan_brace = scan_brace.saturating_sub(1),
b',' if scan_paren == base_paren
&& scan_bracket == base_bracket
&& scan_brace == base_brace =>
{
end = cursor;
break;
}
_ if scan_paren == base_paren
&& scan_bracket == base_bracket
&& scan_brace == base_brace =>
{
if keyword_at(bytes, cursor, b"if") {
found_if = true;
} else if found_if && keyword_at(bytes, cursor, b"else") {
found_else = true;
}
}
_ => {}
}
cursor += 1;
}
let Some(rhs_start) =
(index + 1..end).find(|position| !bytes[*position].is_ascii_whitespace())
else {
continue;
};
let Some(rhs_end) = (index + 1..end)
.rev()
.find(|position| !bytes[*position].is_ascii_whitespace())
.map(|position| position + 1)
else {
continue;
};
if found_if && found_else && !is_fully_parenthesized(bytes, rhs_start, rhs_end)
{
insertions.push((rhs_start, '('));
insertions.push((rhs_end, ')'));
}
}
_ => {}
}
}
if insertions.is_empty() {
return tag.to_owned();
}
insertions.sort_unstable_by_key(|(position, character)| {
(*position, if *character == ')' { 0 } else { 1 })
});
let mut output = String::with_capacity(tag.len() + insertions.len());
let mut insertion_index = 0usize;
for (index, character) in tag.char_indices() {
while insertions
.get(insertion_index)
.is_some_and(|(position, _)| *position == index)
{
output.push(insertions[insertion_index].1);
insertion_index += 1;
}
output.push(character);
}
while insertions
.get(insertion_index)
.is_some_and(|(position, _)| *position == tag.len())
{
output.push(insertions[insertion_index].1);
insertion_index += 1;
}
output
}
let mut output = String::with_capacity(template.len());
let mut cursor = 0usize;
while cursor < template.len() {
let expression = template[cursor..].find("{{").map(|offset| cursor + offset);
let statement = template[cursor..].find("{%").map(|offset| cursor + offset);
let Some(start) = (match (expression, statement) {
(Some(left), Some(right)) => Some(left.min(right)),
(left, right) => left.or(right),
}) else {
output.push_str(&template[cursor..]);
break;
};
output.push_str(&template[cursor..start]);
let delimiter = if template[start..].starts_with("{{") {
"}}"
} else {
"%}"
};
let Some(relative_end) = template[start + 2..].find(delimiter) else {
output.push_str(&template[start..]);
break;
};
let end = start + 2 + relative_end + 2;
output.push_str(&normalize_tag(&template[start..end]));
cursor = end;
}
output
}
fn normalize_chat_template(template: &str) -> String {
normalize_conditional_keyword_arguments(&normalize_generation_blocks(template))
}
fn render_jinja_template<'a, 'defaults, R, T>(
template: Template,
conversations: impl IntoIterator<Item = Chat<'a, R, T>>,
tools: Option<&'a [serde_json::Value]>,
documents: Option<&'a [Document]>,
add_generation_prompt: Option<bool>,
continue_final_message: Option<bool>,
template_kwargs: TemplateKwargs<'defaults, 'a>,
) -> Result<Vec<String>, Error>
where
R: Serialize + 'a,
T: Serialize + 'a,
{
let add_generation_prompt = add_generation_prompt.unwrap_or(false);
let continue_final_message = continue_final_message.unwrap_or(false);
let mut rendered = Vec::new();
for chat in conversations {
let empty_tools: &[serde_json::Value] = &[];
let empty_documents: &[Document] = &[];
let tools = tools.unwrap_or(empty_tools);
let documents = documents.unwrap_or(empty_documents);
let messages = if chat.len() == 1 {
serde_json::to_value(&chat[0].content)
.ok()
.and_then(|value| match value {
serde_json::Value::Array(messages)
if serde_json::to_value(&chat[0].role).ok()
== Some(serde_json::Value::Null) =>
{
Some(serde_json::Value::Array(messages))
}
_ => None,
})
} else {
None
}
.unwrap_or_else(|| serde_json::to_value(&chat).unwrap_or(serde_json::Value::Null));
let mut context = serde_json::Map::new();
context.insert("messages".to_string(), messages);
context.insert("tools".to_string(), serde_json::to_value(tools)?);
context.insert("documents".to_string(), serde_json::to_value(documents)?);
context.insert(
"add_generation_prompt".to_string(),
serde_json::Value::Bool(add_generation_prompt),
);
if let Some(default_template_kwargs) = template_kwargs.defaults {
context.extend(default_template_kwargs.clone());
}
if let Some(template_kwargs) = template_kwargs.overrides {
context.extend(template_kwargs.clone());
}
let mut rendered_chat = template.render(context)?;
rendered_chat = rendered_chat.trim_start_matches('\n').to_string();
if continue_final_message {
let Some(final_message) = chat
.last()
.and_then(|chat| serde_json::to_value(&chat.content).ok())
.and_then(|value| match value {
serde_json::Value::String(text) => Some(text),
other => other
.get("text")
.or_else(|| other.get("content"))
.and_then(|value| value.as_str())
.map(ToString::to_string),
})
else {
continue;
};
let final_message_str = final_message;
if !rendered_chat.contains(final_message_str.trim()) {
return Err(Error::FinalMsgNotInChat);
}
let final_msg_loc = rendered_chat.rfind(&final_message_str.trim()).unwrap();
let final_msg_len = final_message_str.trim_start().len();
rendered_chat = if rendered_chat[final_msg_loc..final_msg_loc + final_msg_len]
== final_message_str
{
rendered_chat[..final_msg_loc + final_msg_len].to_string()
} else {
rendered_chat[..final_msg_loc + final_message_str.trim().len()].to_string()
};
}
rendered.push(rendered_chat);
}
Ok(rendered)
}
#[cfg(test)]
mod tests {
use minijinja::Environment;
use std::{collections::BTreeSet, path::PathBuf};
use crate::tokenizer::{
apply_chat_template, apply_chat_template_json, load_model_chat_template_from_file,
load_model_chat_template_from_str, normalize_conditional_keyword_arguments,
normalize_generation_blocks, ApplyChatTemplateArgs, ChatTemplateIdentity, Conversation,
ModelChatTemplate, Role, Tokenizer,
};
fn fixtures_dir() -> PathBuf {
std::env::var("TEST_MODEL_DIR")
.map(PathBuf::from)
.unwrap_or_else(|_| {
PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("tests/fixtures/qwen3")
})
}
#[test]
fn generation_blocks_become_transparent_minijinja_blocks() {
assert_eq!(
normalize_generation_blocks(
"before{%- generation -%}assistant{%- endgeneration -%}after"
),
"before{%- if true -%}assistant{%- endif -%}after"
);
assert_eq!(
normalize_generation_blocks("{% generation %}{{ generation }}{% endgeneration %}"),
"{% if true %}{{ generation }}{% endif %}"
);
}
#[test]
fn generation_blocks_are_transparent_to_template_kwarg_analysis() {
let kwargs = super::chat_template_kwargs(
concat!(
"{% generation %}",
"{{ messages }}{{ custom_flag }}",
"{% endgeneration %}"
),
"generation-block-template",
)
.unwrap();
assert_eq!(kwargs, BTreeSet::from(["custom_flag".to_owned()]));
}
#[test]
fn conditional_keyword_arguments_are_parenthesized_for_minijinja() {
let template = concat!(
"plain namespace(name=value if condition else 'fallback')",
"{{ namespace(name=value if condition else 'fallback', keep=other) }}",
"{% set wrapped = namespace(name=(value if condition else 'fallback')) %}",
);
let normalized = normalize_conditional_keyword_arguments(template);
assert_eq!(
normalized,
concat!(
"plain namespace(name=value if condition else 'fallback')",
"{{ namespace(name=(value if condition else 'fallback'), keep=other) }}",
"{% set wrapped = namespace(name=(value if condition else 'fallback')) %}",
)
);
let mut env = Environment::new();
env.add_template_owned("conditional-kwarg", normalized)
.unwrap();
}
#[test]
fn released_muse_glimmer_template_renders_reasoning_and_atem_history() {
let template =
include_str!("../tests/fixtures/chat_templates/muse-glimmer-30b-97c77dff.jinja")
.strip_suffix('\n')
.expect("the fixture-only file terminator is documented");
let mut tokenizer = Tokenizer::from_tokenizer(tokenizers::Tokenizer::new(
tokenizers::models::wordlevel::WordLevel::default(),
));
tokenizer.set_template_kwargs(serde_json::Map::from_iter([(
"bos_token".into(),
serde_json::json!(""),
)]));
let kwargs =
serde_json::Map::from_iter([("reasoning_strength".into(), serde_json::json!("xhigh"))]);
let rendered = tokenizer
.apply_chat_template_json(
ModelChatTemplate::Single(template.into()),
[vec![
serde_json::json!({"role": "user", "content": "probe"}),
]],
Some(&[]),
"meta-models/Muse-Glimmer-30B",
true,
Some(&kwargs),
)
.unwrap()
.remove(0);
assert!(rendered.contains("Reasoning strength: xhigh."));
assert!(rendered.ends_with("<|start|>assistant"));
let tools = vec![serde_json::json!({
"type": "function",
"function": {
"name": "lookup",
"description": "look up a value",
"parameters": {
"type": "object",
"properties": {"value": {"type": "string"}}
}
}
})];
let messages = vec![
serde_json::json!({"role": "user", "content": "probe"}),
serde_json::json!({
"role": "assistant",
"reasoning_content": "inspect",
"content": "",
"tool_calls": [{
"id": "call_1",
"type": "function",
"function": {
"name": "lookup",
"arguments": {"value": "probe-value"}
}
}]
}),
serde_json::json!({
"role": "tool",
"name": "lookup",
"tool_call_id": "call_1",
"content": "result"
}),
serde_json::json!({"role": "assistant", "content": "done"}),
];
let history = tokenizer
.apply_chat_template_json(
ModelChatTemplate::Single(template.into()),
[messages],
Some(&tools),
"meta-models/Muse-Glimmer-30B",
false,
None,
)
.unwrap()
.remove(0);
assert!(history.contains(concat!(
"<|start|>assistant to=self<|message|>inspect<|eom|>",
"<|start|>assistant to=lookup<|message|>",
"<atem:function_calls>\n<atem:invoke name=\"lookup\">\n",
"<atem:parameter name=\"value\">probe-value</atem:parameter>"
)));
assert!(history
.contains("<|start|>tool lookup<|message|><tool_output name=\"lookup\">\nresult"));
}
#[test]
fn test_load_chat_template_from_file() {
let file = fixtures_dir().join("tokenizer_config.json");
let chat_template = load_model_chat_template_from_file(file).unwrap().unwrap();
assert!(!chat_template.select(None).unwrap().template().is_empty());
}
#[test]
fn single_chat_template_remains_compatible() {
let templates = load_model_chat_template_from_str(r#"{"chat_template":"single-template"}"#)
.unwrap()
.unwrap();
let selected = templates.select(None).unwrap();
assert_eq!(selected.template(), "single-template");
assert_eq!(selected.identity(), &ChatTemplateIdentity::Single);
let selected_with_tools = templates.select(Some(&[serde_json::json!({})])).unwrap();
assert_eq!(selected_with_tools.template(), "single-template");
assert_eq!(
selected_with_tools.identity(),
&ChatTemplateIdentity::Single
);
}
#[test]
fn named_chat_templates_select_default_or_tool_use() {
let templates = load_model_chat_template_from_str(
r#"{
"chat_template": [
{"name": "tool_use", "template": "tools-template"},
{"name": "default", "template": "default-template"}
]
}"#,
)
.unwrap()
.unwrap();
let selected = templates.select(None).unwrap();
assert_eq!(selected.template(), "default-template");
assert_eq!(
selected.identity(),
&ChatTemplateIdentity::Named("default".into())
);
let selected = templates.select(Some(&[])).unwrap();
assert_eq!(selected.template(), "default-template");
assert_eq!(
selected.identity(),
&ChatTemplateIdentity::Named("default".into())
);
let tools = [serde_json::json!({"type": "function"})];
let selected = templates.select(Some(&tools)).unwrap();
assert_eq!(selected.template(), "tools-template");
assert_eq!(
selected.identity(),
&ChatTemplateIdentity::Named("tool_use".into())
);
let default_only = load_model_chat_template_from_str(
r#"{"chat_template":[{"name":"default","template":"default-only"}]}"#,
)
.unwrap()
.unwrap();
let selected = default_only.select(Some(&tools)).unwrap();
assert_eq!(selected.template(), "default-only");
assert_eq!(
selected.identity(),
&ChatTemplateIdentity::Named("default".into())
);
}
#[test]
fn named_chat_template_selection_reports_missing_default_deterministically() {
let templates = load_model_chat_template_from_str(
r#"{"chat_template":[
{"name":"rag","template":"rag-template"},
{"name":"chat","template":"chat-template"}
]}"#,
)
.unwrap()
.unwrap();
assert_eq!(
templates.select(None).unwrap_err().to_string(),
r#"chat_template collection has no default template; available templates: ["chat", "rag"]"#
);
}
#[test]
fn named_chat_template_collection_rejects_malformed_and_duplicate_entries() {
let cases = [
(
r#"{"chat_template":{}}"#,
"invalid chat_template: expected a string or an array of named template entries",
),
(
r#"{"chat_template":[]}"#,
"invalid chat_template: named template collection must not be empty",
),
(
r#"{"chat_template":["default"]}"#,
"invalid chat_template: entry 0 must be an object with string fields \"name\" and \"template\"",
),
(
r#"{"chat_template":[{"template":"x"}]}"#,
"invalid chat_template: entry 0 must contain exactly the fields \"name\" and \"template\"",
),
(
r#"{"chat_template":[{"name":"default","template":1}]}"#,
"invalid chat_template: entry 0 field \"template\" must be a string",
),
(
r#"{"chat_template":[{"name":"default","template":"a"},{"name":"default","template":"b"}]}"#,
"invalid chat_template: duplicate named template \"default\"",
),
];
for (config, expected) in cases {
assert_eq!(
load_model_chat_template_from_str(config)
.unwrap_err()
.to_string(),
expected
);
}
}
#[test]
fn apply_chat_template_apis_share_named_template_selection() {
let templates = load_model_chat_template_from_str(
r#"{
"chat_template": [
{"name": "default", "template": "default"},
{"name": "tool_use", "template": "tool_use"}
]
}"#,
)
.unwrap()
.unwrap();
let mut env = Environment::new();
let conversations = [vec![Conversation {
role: Role::User,
content: "hello",
}]
.into()];
let rendered = apply_chat_template(
&mut env,
templates.clone(),
ApplyChatTemplateArgs {
conversations,
tools: Some(&[]),
documents: None,
model_id: "selection-test",
chat_template_id: None,
add_generation_prompt: None,
continue_final_message: None,
template_kwargs: None,
},
)
.unwrap();
assert_eq!(rendered, vec!["default"]);
let tools = [serde_json::json!({"type": "function"})];
let rendered = apply_chat_template_json(
&mut env,
templates,
[vec![
serde_json::json!({"role": "user", "content": "hello"}),
]],
Some(&tools),
"selection-test",
false,
None,
)
.unwrap();
assert_eq!(rendered, vec!["tool_use"]);
}
#[test]
fn tokenizer_tojson_accepts_hugging_face_sort_and_separator_kwargs() {
let raw = tokenizers::Tokenizer::new(tokenizers::models::wordlevel::WordLevel::default());
let mut tokenizer = super::Tokenizer::from_tokenizer(raw);
let tools = [serde_json::json!({"zeta": 2, "alpha": 1})];
let rendered = tokenizer
.apply_chat_template_json(
"{{ tools[0] | tojson(sort_keys=true, separators=(\",\", \":\")) }}",
[Vec::new()],
Some(&tools),
"tojson-kwargs-test",
false,
None,
)
.unwrap();
assert_eq!(rendered, [r#"{"alpha":1,"zeta":2}"#]);
}
#[test]
fn test_apply_chat_template() {
let file = fixtures_dir().join("tokenizer_config.json");
let model_chat_template = load_model_chat_template_from_file(file).unwrap().unwrap();
assert!(!model_chat_template
.select(None)
.unwrap()
.template()
.is_empty());
let model_id = "mlx-community/Qwen3-4B-bf16".to_string();
let conversations = vec![Conversation {
role: Role::User,
content: "hello",
}];
let args = ApplyChatTemplateArgs {
conversations: [conversations.into()],
tools: None,
documents: None,
model_id: &model_id,
chat_template_id: None,
add_generation_prompt: None,
continue_final_message: None,
template_kwargs: None,
};
let mut env = Environment::new();
env.set_unknown_method_callback(minijinja_contrib::pycompat::unknown_method_callback);
let rendered_chat = apply_chat_template(&mut env, model_chat_template, args).unwrap();
println!("{:?}", rendered_chat);
}
#[test]
fn test_apply_chat_template_with_template_kwargs() {
let model_template =
"{% if enable_thinking %}think{% else %}no-think{% endif %}".to_string();
let model_id = "test-model".to_string();
let conversations = vec![Conversation {
role: Role::User,
content: "hello",
}];
let mut template_kwargs = serde_json::Map::new();
template_kwargs.insert(
"enable_thinking".to_string(),
serde_json::Value::Bool(false),
);
let args = ApplyChatTemplateArgs {
conversations: [conversations.into()],
tools: None,
documents: None,
model_id: &model_id,
chat_template_id: None,
add_generation_prompt: None,
continue_final_message: None,
template_kwargs: Some(&template_kwargs),
};
let mut env = Environment::new();
env.set_unknown_method_callback(minijinja_contrib::pycompat::unknown_method_callback);
let rendered_chat = apply_chat_template(&mut env, model_template, args).unwrap();
assert_eq!(rendered_chat, vec!["no-think"]);
}
#[test]
fn test_tokenizer_template_kwargs_defaults_can_be_overridden() {
let raw = tokenizers::Tokenizer::new(tokenizers::models::wordlevel::WordLevel::default());
let mut tokenizer = super::Tokenizer::from_tokenizer(raw);
tokenizer.set_template_kwargs(serde_json::Map::from_iter([
(
"bos_token".to_string(),
serde_json::Value::String("<bos>".to_string()),
),
(
"tone".to_string(),
serde_json::Value::String("default".to_string()),
),
]));
let model_template =
"{{ bos_token }}{{ messages[0].role }}: {{ messages[0].content }} {{ tone }}"
.to_string();
let model_id = "test-model".to_string();
let conversations = vec![Conversation {
role: Role::User,
content: "hello",
}];
let mut template_kwargs = serde_json::Map::new();
template_kwargs.insert(
"tone".to_string(),
serde_json::Value::String("override".to_string()),
);
let rendered_chat = tokenizer
.apply_chat_template(
model_template,
ApplyChatTemplateArgs {
conversations: [conversations.into()],
tools: None,
documents: None,
model_id: &model_id,
chat_template_id: None,
add_generation_prompt: None,
continue_final_message: None,
template_kwargs: Some(&template_kwargs),
},
)
.unwrap();
assert_eq!(rendered_chat, vec!["<bos>user: hello override"]);
}
#[test]
fn test_chat_template_kwargs_filters_standard_variables_and_globals() {
let model_template = concat!(
"{% set ns = namespace(found=false) %}",
"{% for message in messages %}{{ message.role }}{% endfor %}",
"{% if tools %}{{ tools|length }}{% endif %}",
"{% if documents %}{{ documents|length }}{% endif %}",
"{% if add_generation_prompt and enable_thinking is defined %}{{ tone }}{% endif %}",
);
let kwargs = super::chat_template_kwargs(model_template, "test-model").unwrap();
assert_eq!(
kwargs.into_iter().collect::<Vec<_>>(),
vec!["enable_thinking", "tone"]
);
}
#[test]
fn test_qwen_fixture_reports_enable_thinking_kwarg() {
let file = fixtures_dir().join("tokenizer_config.json");
let chat_template = load_model_chat_template_from_file(file).unwrap().unwrap();
let kwargs = super::chat_template_kwargs(
chat_template.select(None).unwrap().template(),
"qwen-fixture",
)
.unwrap();
assert!(kwargs.contains("enable_thinking"), "{kwargs:?}");
}
#[test]
#[ignore = "requires local model files (tokenizer.json is 11MB)"]
fn test_tokenizer_apply_chat_template() {
let tokenizer_file = fixtures_dir().join("tokenizer.json");
let tokenizer_config_file = fixtures_dir().join("tokenizer_config.json");
let model_id = "mlx-community/Qwen3-4B-bf16".to_string();
let conversations = vec![Conversation {
role: Role::User,
content: "hello",
}];
let mut tokenizer = super::Tokenizer::from_file(tokenizer_file).unwrap();
let model_chat_template = load_model_chat_template_from_file(tokenizer_config_file)
.unwrap()
.unwrap();
assert!(!model_chat_template
.select(None)
.unwrap()
.template()
.is_empty());
let args = ApplyChatTemplateArgs {
conversations: [conversations.into()],
tools: None,
documents: None,
model_id: &model_id,
chat_template_id: None,
add_generation_prompt: None,
continue_final_message: None,
template_kwargs: None,
};
let rendered_chat = tokenizer
.apply_chat_template(model_chat_template, args)
.unwrap();
println!("{:?}", rendered_chat);
}
#[test]
#[ignore = "requires local model files (tokenizer.json is 11MB)"]
fn test_tokenizer_apply_chat_template_and_encode() {
let tokenizer_file = fixtures_dir().join("tokenizer.json");
let tokenizer_config_file = fixtures_dir().join("tokenizer_config.json");
let model_id = "mlx-community/Qwen3-4B-bf16".to_string();
let conversations = vec![Conversation {
role: Role::User,
content: "hello",
}];
let mut tokenizer = super::Tokenizer::from_file(tokenizer_file).unwrap();
let model_chat_template = load_model_chat_template_from_file(tokenizer_config_file)
.unwrap()
.unwrap();
assert!(!model_chat_template
.select(None)
.unwrap()
.template()
.is_empty());
let args = ApplyChatTemplateArgs {
conversations: [conversations.into()],
tools: None,
documents: None,
model_id: &model_id,
chat_template_id: None,
add_generation_prompt: None,
continue_final_message: None,
template_kwargs: None,
};
let encodings = tokenizer
.apply_chat_template_and_encode(model_chat_template, args)
.unwrap();
println!("{:?}", encodings.iter().flat_map(|e| e.get_ids()));
}
}