1use std::collections::HashMap;
3
4use crate::prelude::*;
5
6#[derive(Debug)]
8pub struct Template
9{
10 #[allow(dead_code)]
11 source: String,
12 add_generation_prompt: bool,
13}
14
15impl Template
16{
17 pub fn new(source: impl Into<String>) -> Result<Template>
19 {
20 let source = source.into();
21 Ok(Template {
22 source,
23 add_generation_prompt: true,
24 })
25 }
26 #[allow(dead_code)]
27 pub(crate) fn render(
28 &self,
29 messages: &[crate::Message],
30 enable_thinking: bool,
31 options: HashMap<String, minijinja::Value>,
32 ) -> Result<String>
33 {
34 let mut env = minijinja::Environment::new();
35 env.add_template("template", &self.source)?;
36 let rendered = env.get_template("template")?.render(minijinja::context! {
37 messages => messages,
38 add_generation_prompt => self.add_generation_prompt,
39 enable_thinking => enable_thinking,
40 ..options
41 })?;
42 Ok(rendered)
43 }
44}
45
46#[cfg(test)]
47mod tests
48{
49 #[test]
50 fn template_llama()
51 {
52 let template = super::Template::new(include_str!("../data/templates/llama")).unwrap();
53 assert_eq!(
54 template.render(&[crate::Message {
55 role: crate::Role::User,
56 content: "Hello world.".into()
57 }], false, Default::default()).unwrap(),
58 "<|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"
59 );
60 }
61}