Skip to main content

kproc_llm/
template.rs

1//! Template module
2use std::collections::HashMap;
3
4use crate::prelude::*;
5
6/// Template
7#[derive(Debug)]
8pub struct Template
9{
10  #[allow(dead_code)]
11  source: String,
12  add_generation_prompt: bool,
13}
14
15impl Template
16{
17  /// New template
18  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}