use std::collections::HashMap;
use crate::prelude::*;
#[derive(Debug)]
pub struct Template
{
#[allow(dead_code)]
source: String,
add_generation_prompt: bool,
}
impl Template
{
pub fn new(source: impl Into<String>) -> Result<Template>
{
let source = source.into();
Ok(Template {
source,
add_generation_prompt: true,
})
}
#[allow(dead_code)]
pub(crate) fn render(
&self,
messages: &[crate::Message],
enable_thinking: bool,
options: HashMap<String, minijinja::Value>,
) -> Result<String>
{
let mut env = minijinja::Environment::new();
env.add_template("template", &self.source)?;
let rendered = env.get_template("template")?.render(minijinja::context! {
messages => messages,
add_generation_prompt => self.add_generation_prompt,
enable_thinking => enable_thinking,
..options
})?;
Ok(rendered)
}
}
#[cfg(test)]
mod tests
{
#[test]
fn template_llama()
{
let template = super::Template::new(include_str!("../data/templates/llama")).unwrap();
assert_eq!(
template.render(&[crate::Message {
role: crate::Role::User,
content: "Hello world.".into()
}], false, Default::default()).unwrap(),
"<|start_header_id|>user<|end_header_id|>Hello world.<|eot_id|>\n<|start_header_id|>assistant<|end_header_id|><|start_header_id|>assistant<|end_header_id|>\n"
);
}
}