Skip to main content

tokio_prompt_orchestrator/
prompt_template.rs

1//! Jinja-lite template engine for prompt construction.
2//!
3//! Supports `{{ variable }}` substitution with filters, `{% if/elif/else/endif %}` blocks,
4//! `{% for item in list %}...{% endfor %}` loops, and `{# comment #}` removal.
5
6use std::collections::HashMap;
7use std::fmt;
8
9// ── Errors ────────────────────────────────────────────────────────────────────
10
11/// Errors that can occur during template parsing or rendering.
12#[derive(Debug, Clone, PartialEq)]
13pub enum TemplateError {
14    /// A block (`{%`, `{{`, `{#`) was opened but never closed.
15    UnclosedBlock,
16    /// A variable referenced in the template is not present in the context.
17    UnknownVariable(String),
18    /// The template source contains invalid syntax.
19    InvalidSyntax(String),
20    /// Template blocks are nested beyond the allowed limit.
21    NestedBlockLimit,
22}
23
24impl fmt::Display for TemplateError {
25    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
26        match self {
27            TemplateError::UnclosedBlock => write!(f, "unclosed block in template"),
28            TemplateError::UnknownVariable(v) => write!(f, "unknown variable: {v}"),
29            TemplateError::InvalidSyntax(msg) => write!(f, "invalid syntax: {msg}"),
30            TemplateError::NestedBlockLimit => write!(f, "nested block limit exceeded"),
31        }
32    }
33}
34
35impl std::error::Error for TemplateError {}
36
37// ── TemplateVar ───────────────────────────────────────────────────────────────
38
39/// A dynamically-typed value that can be stored in a [`TemplateContext`].
40#[derive(Debug, Clone, PartialEq)]
41pub enum TemplateVar {
42    /// A UTF-8 string.
43    Str(String),
44    /// A 64-bit signed integer.
45    Int(i64),
46    /// A 64-bit floating-point number.
47    Float(f64),
48    /// A boolean.
49    Bool(bool),
50    /// An ordered list of [`TemplateVar`] values.
51    List(Vec<TemplateVar>),
52    /// A string-keyed map of [`TemplateVar`] values.
53    Map(HashMap<String, TemplateVar>),
54}
55
56impl TemplateVar {
57    /// Convert this value to its string representation.
58    pub fn as_str(&self) -> String {
59        match self {
60            TemplateVar::Str(s) => s.clone(),
61            TemplateVar::Int(i) => i.to_string(),
62            TemplateVar::Float(f) => f.to_string(),
63            TemplateVar::Bool(b) => b.to_string(),
64            TemplateVar::List(items) => {
65                let parts: Vec<String> = items.iter().map(|v| v.as_str()).collect();
66                format!("[{}]", parts.join(", "))
67            }
68            TemplateVar::Map(m) => {
69                let parts: Vec<String> = m.iter().map(|(k, v)| format!("{k}: {}", v.as_str())).collect();
70                format!("{{{}}}", parts.join(", "))
71            }
72        }
73    }
74
75    /// Convert this value to bool (for conditionals).
76    pub fn as_bool(&self) -> bool {
77        self.is_truthy()
78    }
79
80    /// Whether this value is considered truthy.
81    pub fn is_truthy(&self) -> bool {
82        match self {
83            TemplateVar::Bool(b) => *b,
84            TemplateVar::Int(i) => *i != 0,
85            TemplateVar::Float(f) => *f != 0.0,
86            TemplateVar::Str(s) => !s.is_empty(),
87            TemplateVar::List(l) => !l.is_empty(),
88            TemplateVar::Map(m) => !m.is_empty(),
89        }
90    }
91}
92
93// ── TemplateContext ───────────────────────────────────────────────────────────
94
95/// A variable binding context passed to [`Template::render`].
96#[derive(Debug, Clone, Default)]
97pub struct TemplateContext {
98    vars: HashMap<String, TemplateVar>,
99}
100
101impl TemplateContext {
102    /// Create an empty context.
103    pub fn new() -> Self {
104        Self::default()
105    }
106
107    /// Insert any [`TemplateVar`] under `key`.
108    pub fn set(&mut self, key: &str, val: TemplateVar) {
109        self.vars.insert(key.to_string(), val);
110    }
111
112    /// Retrieve a variable by key.
113    pub fn get(&self, key: &str) -> Option<&TemplateVar> {
114        self.vars.get(key)
115    }
116
117    /// Convenience: insert a string variable.
118    pub fn set_str(&mut self, key: &str, val: &str) {
119        self.set(key, TemplateVar::Str(val.to_string()));
120    }
121
122    /// Convenience: insert an integer variable.
123    pub fn set_int(&mut self, key: &str, val: i64) {
124        self.set(key, TemplateVar::Int(val));
125    }
126
127    /// Convenience: insert a boolean variable.
128    pub fn set_bool(&mut self, key: &str, val: bool) {
129        self.set(key, TemplateVar::Bool(val));
130    }
131
132    /// Convenience: insert a list variable.
133    pub fn set_list(&mut self, key: &str, val: Vec<TemplateVar>) {
134        self.set(key, TemplateVar::List(val));
135    }
136}
137
138// ── Internal AST ─────────────────────────────────────────────────────────────
139
140/// Internal pre-parsed node.
141#[derive(Debug, Clone)]
142enum Node {
143    /// Literal text, emitted verbatim.
144    Text(String),
145    /// `{{ expr }}` — variable substitution, possibly with a filter chain.
146    Expr(String),
147    /// `{% if cond %}` block. Contains (condition, body, optional elif/else chain).
148    If {
149        branches: Vec<(String, Vec<Node>)>, // (condition, nodes)
150        else_branch: Option<Vec<Node>>,
151    },
152    /// `{% for item in list %}` loop.
153    For {
154        item: String,
155        list: String,
156        body: Vec<Node>,
157    },
158}
159
160// ── Template ──────────────────────────────────────────────────────────────────
161
162/// A pre-parsed template that can be rendered repeatedly with different contexts.
163#[derive(Debug, Clone)]
164pub struct Template {
165    source: String,
166    nodes: Vec<Node>,
167}
168
169const MAX_NEST: usize = 32;
170
171impl Template {
172    /// Parse `source` into a [`Template`].
173    pub fn new(source: &str) -> Result<Self, TemplateError> {
174        let nodes = parse(source)?;
175        Ok(Self { source: source.to_string(), nodes })
176    }
177
178    /// Render this template with the given context.
179    pub fn render(&self, ctx: &TemplateContext) -> Result<String, TemplateError> {
180        render_nodes(&self.nodes, ctx)
181    }
182
183    /// Return all variable names referenced in `{{ ... }}` expressions.
184    pub fn variables(&self) -> Vec<String> {
185        let mut out = Vec::new();
186        collect_variables(&self.nodes, &mut out);
187        out.sort();
188        out.dedup();
189        out
190    }
191
192    /// Validate that all blocks are properly paired without rendering.
193    pub fn validate(&self) -> Result<(), TemplateError> {
194        // Parsing already validates block structure; re-parse to confirm.
195        parse(&self.source).map(|_| ())
196    }
197}
198
199fn collect_variables(nodes: &[Node], out: &mut Vec<String>) {
200    for node in nodes {
201        match node {
202            Node::Text(_) => {}
203            Node::Expr(expr) => {
204                let var = expr.split('|').next().unwrap_or("").trim().to_string();
205                // handle dot-access: take root
206                let root = var.split('.').next().unwrap_or("").trim().to_string();
207                if !root.is_empty() && !root.starts_with("loop.") {
208                    out.push(root);
209                }
210            }
211            Node::If { branches, else_branch } => {
212                for (cond, body) in branches {
213                    let root = cond.trim().split('.').next().unwrap_or("").trim().to_string();
214                    if !root.is_empty() {
215                        out.push(root);
216                    }
217                    collect_variables(body, out);
218                }
219                if let Some(eb) = else_branch {
220                    collect_variables(eb, out);
221                }
222            }
223            Node::For { list, body, .. } => {
224                out.push(list.trim().to_string());
225                collect_variables(body, out);
226            }
227        }
228    }
229}
230
231// ── Parser ────────────────────────────────────────────────────────────────────
232
233/// Token produced by the lexer.
234#[derive(Debug, Clone)]
235enum Token {
236    Text(String),
237    Expr(String),   // {{ ... }}
238    Tag(String),    // {% ... %}
239    Comment,        // {# ... #}
240}
241
242fn lex(source: &str) -> Result<Vec<Token>, TemplateError> {
243    let mut tokens = Vec::new();
244    let mut chars: &str = source;
245
246    while !chars.is_empty() {
247        if chars.starts_with("{{") {
248            let end = chars.find("}}").ok_or(TemplateError::UnclosedBlock)?;
249            let inner = chars[2..end].trim().to_string();
250            tokens.push(Token::Expr(inner));
251            chars = &chars[end + 2..];
252        } else if chars.starts_with("{%") {
253            let end = chars.find("%}").ok_or(TemplateError::UnclosedBlock)?;
254            let inner = chars[2..end].trim().to_string();
255            tokens.push(Token::Tag(inner));
256            chars = &chars[end + 2..];
257        } else if chars.starts_with("{#") {
258            let end = chars.find("#}").ok_or(TemplateError::UnclosedBlock)?;
259            tokens.push(Token::Comment);
260            chars = &chars[end + 2..];
261        } else {
262            // Find the next `{{`, `{%`, or `{#`
263            let next = chars.find("{%")
264                .unwrap_or(chars.len())
265                .min(chars.find("{{").unwrap_or(chars.len()))
266                .min(chars.find("{#").unwrap_or(chars.len()));
267            tokens.push(Token::Text(chars[..next].to_string()));
268            chars = &chars[next..];
269        }
270    }
271
272    Ok(tokens)
273}
274
275/// Recursive descent parser — returns a list of top-level nodes.
276fn parse(source: &str) -> Result<Vec<Node>, TemplateError> {
277    let tokens = lex(source)?;
278    let mut pos = 0;
279    let nodes = parse_nodes(&tokens, &mut pos, 0)?;
280    Ok(nodes)
281}
282
283fn parse_nodes(tokens: &[Token], pos: &mut usize, depth: usize) -> Result<Vec<Node>, TemplateError> {
284    if depth > MAX_NEST {
285        return Err(TemplateError::NestedBlockLimit);
286    }
287    let mut nodes = Vec::new();
288
289    while *pos < tokens.len() {
290        match &tokens[*pos] {
291            Token::Text(t) => {
292                nodes.push(Node::Text(t.clone()));
293                *pos += 1;
294            }
295            Token::Expr(e) => {
296                nodes.push(Node::Expr(e.clone()));
297                *pos += 1;
298            }
299            Token::Comment => {
300                *pos += 1;
301            }
302            Token::Tag(tag) => {
303                let t = tag.as_str();
304                if t == "endif" || t == "endfor" || t.starts_with("else") || t.starts_with("elif") {
305                    // Stop: caller handles these.
306                    break;
307                } else if t.starts_with("if ") {
308                    *pos += 1;
309                    let node = parse_if(tokens, pos, t, depth)?;
310                    nodes.push(node);
311                } else if t.starts_with("for ") {
312                    *pos += 1;
313                    let node = parse_for(tokens, pos, t, depth)?;
314                    nodes.push(node);
315                } else {
316                    return Err(TemplateError::InvalidSyntax(format!("unknown tag: {t}")));
317                }
318            }
319        }
320    }
321
322    Ok(nodes)
323}
324
325fn parse_if(tokens: &[Token], pos: &mut usize, first_tag: &str, depth: usize) -> Result<Node, TemplateError> {
326    // first_tag = "if <cond>"
327    let first_cond = first_tag["if ".len()..].trim().to_string();
328    let first_body = parse_nodes(tokens, pos, depth + 1)?;
329
330    let mut branches: Vec<(String, Vec<Node>)> = vec![(first_cond, first_body)];
331    let mut else_branch: Option<Vec<Node>> = None;
332
333    loop {
334        if *pos >= tokens.len() {
335            return Err(TemplateError::UnclosedBlock);
336        }
337        match &tokens[*pos] {
338            Token::Tag(tag) => {
339                let t = tag.as_str();
340                if t == "endif" {
341                    *pos += 1;
342                    break;
343                } else if let Some(rest) = t.strip_prefix("elif ") {
344                    let cond = rest.trim().to_string();
345                    *pos += 1;
346                    let body = parse_nodes(tokens, pos, depth + 1)?;
347                    branches.push((cond, body));
348                } else if t == "else" {
349                    *pos += 1;
350                    let body = parse_nodes(tokens, pos, depth + 1)?;
351                    else_branch = Some(body);
352                } else {
353                    return Err(TemplateError::InvalidSyntax(format!("unexpected tag in if: {t}")));
354                }
355            }
356            _ => return Err(TemplateError::UnclosedBlock),
357        }
358    }
359
360    Ok(Node::If { branches, else_branch })
361}
362
363fn parse_for(tokens: &[Token], pos: &mut usize, tag: &str, depth: usize) -> Result<Node, TemplateError> {
364    // tag = "for <item> in <list>"
365    let rest = &tag["for ".len()..];
366    let parts: Vec<&str> = rest.splitn(3, ' ').collect();
367    if parts.len() < 3 || parts[1] != "in" {
368        return Err(TemplateError::InvalidSyntax(format!("malformed for tag: {tag}")));
369    }
370    let item = parts[0].trim().to_string();
371    let list = parts[2].trim().to_string();
372
373    let body = parse_nodes(tokens, pos, depth + 1)?;
374
375    if *pos >= tokens.len() {
376        return Err(TemplateError::UnclosedBlock);
377    }
378    match &tokens[*pos] {
379        Token::Tag(t) if t == "endfor" => {
380            *pos += 1;
381        }
382        _ => return Err(TemplateError::UnclosedBlock),
383    }
384
385    Ok(Node::For { item, list, body })
386}
387
388// ── Renderer ──────────────────────────────────────────────────────────────────
389
390fn render_nodes(nodes: &[Node], ctx: &TemplateContext) -> Result<String, TemplateError> {
391    let mut out = String::new();
392    for node in nodes {
393        match node {
394            Node::Text(t) => out.push_str(t),
395            Node::Expr(expr) => {
396                out.push_str(&eval_expr(expr, ctx)?);
397            }
398            Node::If { branches, else_branch } => {
399                let mut matched = false;
400                for (cond, body) in branches {
401                    if eval_condition(cond, ctx)? {
402                        out.push_str(&render_nodes(body, ctx)?);
403                        matched = true;
404                        break;
405                    }
406                }
407                if !matched {
408                    if let Some(eb) = else_branch {
409                        out.push_str(&render_nodes(eb, ctx)?);
410                    }
411                }
412            }
413            Node::For { item, list, body } => {
414                let list_var = ctx.get(list).ok_or_else(|| TemplateError::UnknownVariable(list.clone()))?;
415                let items = match list_var {
416                    TemplateVar::List(l) => l.clone(),
417                    other => vec![other.clone()],
418                };
419                let len = items.len();
420                for (i, val) in items.into_iter().enumerate() {
421                    let mut inner_ctx = ctx.clone();
422                    inner_ctx.set(item, val);
423                    inner_ctx.set("loop.index", TemplateVar::Int((i + 1) as i64));
424                    inner_ctx.set("loop.index0", TemplateVar::Int(i as i64));
425                    inner_ctx.set("loop.first", TemplateVar::Bool(i == 0));
426                    inner_ctx.set("loop.last", TemplateVar::Bool(i == len - 1));
427                    out.push_str(&render_nodes(body, &inner_ctx)?);
428                }
429            }
430        }
431    }
432    Ok(out)
433}
434
435/// Evaluate a `{{ expr }}` expression (variable lookup + optional filters).
436fn eval_expr(expr: &str, ctx: &TemplateContext) -> Result<String, TemplateError> {
437    // Handle loop.* pseudo-variables.
438    if expr.starts_with("loop.") {
439        let key = expr.trim();
440        return match ctx.get(key) {
441            Some(v) => Ok(v.as_str()),
442            None => Ok(String::new()),
443        };
444    }
445
446    let parts: Vec<&str> = expr.splitn(2, '|').collect();
447    let var_name = parts[0].trim();
448
449    // Resolve variable (support dot-path like "obj.field"). A missing
450    // variable is only an error when no `default:` filter can supply a value.
451    let has_default = parts
452        .get(1)
453        .is_some_and(|f| f.split('|').any(|p| p.trim().starts_with("default:")));
454    let value = match resolve_var(var_name, ctx) {
455        Ok(v) => v,
456        Err(_) if has_default => String::new(),
457        Err(e) => return Err(e),
458    };
459
460    // Apply filters if any.
461    if parts.len() == 2 {
462        apply_filter(value, parts[1].trim())
463    } else {
464        Ok(value)
465    }
466}
467
468fn resolve_var(name: &str, ctx: &TemplateContext) -> Result<String, TemplateError> {
469    // Support simple dot-path: "a.b" -> ctx["a"] as Map -> "b"
470    let segments: Vec<&str> = name.splitn(2, '.').collect();
471    let root = segments[0].trim();
472
473    match ctx.get(root) {
474        Some(TemplateVar::Map(m)) if segments.len() == 2 => {
475            let key = segments[1].trim();
476            match m.get(key) {
477                Some(v) => Ok(v.as_str()),
478                None => Err(TemplateError::UnknownVariable(name.to_string())),
479            }
480        }
481        Some(v) => Ok(v.as_str()),
482        None => Err(TemplateError::UnknownVariable(root.to_string())),
483    }
484}
485
486fn apply_filter(value: String, filter_expr: &str) -> Result<String, TemplateError> {
487    // Filters can be chained: `upper | trim`
488    let mut result = value;
489    for part in filter_expr.split('|') {
490        let f = part.trim();
491        result = apply_single_filter(result, f)?;
492    }
493    Ok(result)
494}
495
496fn apply_single_filter(value: String, filter: &str) -> Result<String, TemplateError> {
497    if filter == "upper" {
498        Ok(value.to_uppercase())
499    } else if filter == "lower" {
500        Ok(value.to_lowercase())
501    } else if filter == "trim" {
502        Ok(value.trim().to_string())
503    } else if let Some(rest) = filter.strip_prefix("truncate:") {
504        let n: usize = rest.trim().parse().map_err(|_| TemplateError::InvalidSyntax(format!("truncate: expected number, got {rest}")))?;
505        // Count characters, not bytes, so multi-byte text never splits mid-char.
506        match value.char_indices().nth(n) {
507            Some((cut, _)) => Ok(format!("{}...", &value[..cut])),
508            None => Ok(value),
509        }
510    } else if let Some(rest) = filter.strip_prefix("default:") {
511        if value.is_empty() { Ok(rest.trim().to_string()) } else { Ok(value) }
512    } else {
513        Err(TemplateError::InvalidSyntax(format!("unknown filter: {filter}")))
514    }
515}
516
517fn eval_condition(cond: &str, ctx: &TemplateContext) -> Result<bool, TemplateError> {
518    let cond = cond.trim();
519
520    // Handle "not X"
521    if let Some(inner) = cond.strip_prefix("not ") {
522        return eval_condition(inner.trim(), ctx).map(|b| !b);
523    }
524
525    // Handle "X == Y"
526    if let Some(idx) = cond.find(" == ") {
527        let lhs = cond[..idx].trim();
528        let rhs = cond[idx + 4..].trim().trim_matches('"');
529        let lhs_val = resolve_var(lhs, ctx).unwrap_or_default();
530        return Ok(lhs_val == rhs);
531    }
532
533    // Handle "X != Y"
534    if let Some(idx) = cond.find(" != ") {
535        let lhs = cond[..idx].trim();
536        let rhs = cond[idx + 4..].trim().trim_matches('"');
537        let lhs_val = resolve_var(lhs, ctx).unwrap_or_default();
538        return Ok(lhs_val != rhs);
539    }
540
541    // Simple truthiness check.
542    match ctx.get(cond) {
543        Some(v) => Ok(v.is_truthy()),
544        None => Ok(false),
545    }
546}
547
548// ── TemplateLibrary ───────────────────────────────────────────────────────────
549
550/// A named collection of pre-parsed templates.
551#[derive(Debug, Default)]
552pub struct TemplateLibrary {
553    templates: HashMap<String, Template>,
554}
555
556impl TemplateLibrary {
557    /// Create an empty library.
558    pub fn new() -> Self {
559        Self::default()
560    }
561
562    /// Register a template under `name`, parsing it immediately.
563    pub fn register(&mut self, name: &str, source: &str) -> Result<(), TemplateError> {
564        let tpl = Template::new(source)?;
565        self.templates.insert(name.to_string(), tpl);
566        Ok(())
567    }
568
569    /// Render a named template.
570    pub fn render(&self, name: &str, ctx: &TemplateContext) -> Result<String, TemplateError> {
571        match self.templates.get(name) {
572            Some(t) => t.render(ctx),
573            None => Err(TemplateError::UnknownVariable(name.to_string())),
574        }
575    }
576
577    /// Create a library pre-loaded with common prompt templates.
578    pub fn built_in_templates() -> Self {
579        let mut lib = Self::new();
580
581        lib.register(
582            "chat_completion",
583            r#"{# chat_completion template #}
584{% if system_prompt %}System: {{ system_prompt }}
585
586{% endif %}{% for msg in messages %}{{ msg }}
587{% endfor %}Assistant:"#,
588        ).ok();
589
590        lib.register(
591            "few_shot",
592            r#"{# few_shot template #}
593{{ instruction }}
594
595{% for example in examples %}Example {{ loop.index }}:
596Input: {{ example }}
597Output: [see below]
598
599{% endfor %}Now answer:
600Input: {{ input }}"#,
601        ).ok();
602
603        lib.register(
604            "chain_of_thought",
605            r#"{# chain_of_thought template #}
606{{ problem }}
607
608Let's think step by step.
609{% if context %}
610Context: {{ context }}
611{% endif %}
612Answer:"#,
613        ).ok();
614
615        lib.register(
616            "summarize",
617            r#"{# summarize template #}
618Please summarize the following text{% if max_words %} in at most {{ max_words }} words{% endif %}:
619
620{{ text }}
621
622Summary:"#,
623        ).ok();
624
625        lib.register(
626            "translate",
627            r#"{# translate template #}
628Translate the following text from {{ source_lang | default:English }} to {{ target_lang }}:
629
630{{ text }}
631
632Translation:"#,
633        ).ok();
634
635        lib
636    }
637}
638
639#[cfg(test)]
640mod tests {
641    use super::*;
642
643    #[test]
644    fn test_simple_variable() {
645        let tpl = Template::new("Hello, {{ name }}!").unwrap();
646        let mut ctx = TemplateContext::new();
647        ctx.set_str("name", "World");
648        assert_eq!(tpl.render(&ctx).unwrap(), "Hello, World!");
649    }
650
651    #[test]
652    fn test_filter_upper() {
653        let tpl = Template::new("{{ name | upper }}").unwrap();
654        let mut ctx = TemplateContext::new();
655        ctx.set_str("name", "hello");
656        assert_eq!(tpl.render(&ctx).unwrap(), "HELLO");
657    }
658
659    #[test]
660    fn test_if_else() {
661        let tpl = Template::new("{% if flag %}yes{% else %}no{% endif %}").unwrap();
662        let mut ctx = TemplateContext::new();
663        ctx.set_bool("flag", true);
664        assert_eq!(tpl.render(&ctx).unwrap(), "yes");
665        ctx.set_bool("flag", false);
666        assert_eq!(tpl.render(&ctx).unwrap(), "no");
667    }
668
669    #[test]
670    fn test_for_loop() {
671        let tpl = Template::new("{% for item in items %}{{ item }} {% endfor %}").unwrap();
672        let mut ctx = TemplateContext::new();
673        ctx.set_list("items", vec![
674            TemplateVar::Str("a".into()),
675            TemplateVar::Str("b".into()),
676        ]);
677        assert_eq!(tpl.render(&ctx).unwrap(), "a b ");
678    }
679
680    #[test]
681    fn test_comment_removed() {
682        let tpl = Template::new("before{# ignored #}after").unwrap();
683        let ctx = TemplateContext::new();
684        assert_eq!(tpl.render(&ctx).unwrap(), "beforeafter");
685    }
686
687    #[test]
688    fn test_built_in_templates() {
689        let lib = TemplateLibrary::built_in_templates();
690        let mut ctx = TemplateContext::new();
691        ctx.set_str("text", "The quick brown fox.");
692        ctx.set_str("target_lang", "French");
693        let result = lib.render("translate", &ctx);
694        assert!(result.is_ok(), "{result:?}");
695    }
696}