Skip to main content

tokio_prompt_orchestrator/
template.rs

1//! # Prompt Template Engine
2//!
3//! Simple `{{variable}}` substitution engine with filter support.
4//!
5//! ## Variable Syntax
6//!
7//! - `{{var}}` — plain substitution.
8//! - `{{var | upper}}` — apply the `upper` filter.
9//! - `{{var | lower}}` — apply the `lower` filter.
10//! - `{{var | truncate:100}}` — truncate to 100 characters.
11//! - `{{var | default:"fallback"}}` — use `"fallback"` if `var` is missing.
12//!
13//! ## Example
14//!
15//! ```rust
16//! use tokio_prompt_orchestrator::template::{
17//!     PromptTemplate, TemplateContext, TemplateValue,
18//! };
19//!
20//! let t = PromptTemplate::from_str("Hello, {{name | upper}}!").unwrap();
21//! let mut ctx = TemplateContext::new();
22//! ctx.set("name", TemplateValue::Text("world".to_string()));
23//! let rendered = t.render(&ctx).unwrap();
24//! assert_eq!(rendered, "Hello, WORLD!");
25//! ```
26
27use std::collections::HashMap;
28use thiserror::Error;
29
30// ── Errors ────────────────────────────────────────────────────────────────────
31
32/// Errors that can occur during template parsing or rendering.
33#[derive(Debug, Error, PartialEq, Clone)]
34pub enum TemplateError {
35    /// A `{{variable}}` referenced in the template has no value in the context
36    /// and no `default` filter was supplied.
37    #[error("missing variable: {0}")]
38    MissingVariable(String),
39
40    /// An unrecognised filter name was used (e.g. `{{x | frobnicate}}`).
41    #[error("invalid filter: {0}")]
42    InvalidFilter(String),
43
44    /// The template text could not be parsed (e.g. an unclosed `{{`).
45    #[error("parse error: {0}")]
46    ParseError(String),
47}
48
49// ── TemplateValue ─────────────────────────────────────────────────────────────
50
51/// A typed value that can be stored in a [`TemplateContext`].
52#[derive(Debug, Clone, PartialEq)]
53pub enum TemplateValue {
54    /// A plain text string.
55    Text(String),
56    /// A number (rendered as its default decimal representation).
57    Number(f64),
58    /// A boolean (rendered as `"true"` or `"false"`).
59    Bool(bool),
60    /// A list of strings (rendered as comma-separated).
61    List(Vec<String>),
62}
63
64impl TemplateValue {
65    /// Convert the value to its string representation.
66    pub fn to_string_repr(&self) -> String {
67        match self {
68            TemplateValue::Text(s) => s.clone(),
69            TemplateValue::Number(n) => {
70                if n.fract() == 0.0 && n.abs() < 1e15_f64 {
71                    format!("{}", *n as i64)
72                } else {
73                    format!("{n}")
74                }
75            }
76            TemplateValue::Bool(b) => b.to_string(),
77            TemplateValue::List(items) => items.join(", "),
78        }
79    }
80}
81
82// ── TemplateContext ───────────────────────────────────────────────────────────
83
84/// Variable bindings for template rendering.
85#[derive(Debug, Clone, Default)]
86pub struct TemplateContext {
87    values: HashMap<String, TemplateValue>,
88}
89
90impl TemplateContext {
91    /// Create an empty context.
92    pub fn new() -> Self {
93        Self::default()
94    }
95
96    /// Insert a variable binding.
97    pub fn set(&mut self, key: impl Into<String>, value: TemplateValue) {
98        self.values.insert(key.into(), value);
99    }
100
101    /// Retrieve a variable by name.
102    pub fn get(&self, key: &str) -> Option<&TemplateValue> {
103        self.values.get(key)
104    }
105}
106
107// ── Filter ────────────────────────────────────────────────────────────────────
108
109/// A filter applied to a substituted value.
110#[derive(Debug, Clone, PartialEq)]
111enum Filter {
112    Upper,
113    Lower,
114    Truncate(usize),
115    Default(String),
116}
117
118impl Filter {
119    fn parse(spec: &str) -> Result<Self, TemplateError> {
120        let spec = spec.trim();
121        if spec == "upper" {
122            return Ok(Filter::Upper);
123        }
124        if spec == "lower" {
125            return Ok(Filter::Lower);
126        }
127        if let Some(arg) = spec.strip_prefix("truncate:") {
128            let n: usize = arg.trim().parse().map_err(|_| {
129                TemplateError::InvalidFilter(format!("truncate: invalid length '{arg}'"))
130            })?;
131            return Ok(Filter::Truncate(n));
132        }
133        if let Some(arg) = spec.strip_prefix("default:") {
134            let arg = arg.trim();
135            // Strip surrounding quotes if present.
136            let value = if (arg.starts_with('"') && arg.ends_with('"'))
137                || (arg.starts_with('\'') && arg.ends_with('\''))
138            {
139                arg[1..arg.len() - 1].to_string()
140            } else {
141                arg.to_string()
142            };
143            return Ok(Filter::Default(value));
144        }
145        Err(TemplateError::InvalidFilter(spec.to_string()))
146    }
147
148    fn apply(&self, value: &str) -> String {
149        match self {
150            Filter::Upper => value.to_uppercase(),
151            Filter::Lower => value.to_lowercase(),
152            Filter::Truncate(n) => {
153                if value.len() <= *n {
154                    value.to_string()
155                } else {
156                    value[..*n].to_string()
157                }
158            }
159            Filter::Default(_) => value.to_string(), // used only when value is missing
160        }
161    }
162}
163
164// ── Token ─────────────────────────────────────────────────────────────────────
165
166#[derive(Debug, Clone)]
167enum Token {
168    Literal(String),
169    Substitution {
170        variable: String,
171        filters: Vec<Filter>,
172    },
173}
174
175// ── PromptTemplate ────────────────────────────────────────────────────────────
176
177/// A compiled prompt template.
178///
179/// Parse once with [`PromptTemplate::from_str`], render many times with
180/// [`PromptTemplate::render`].
181#[derive(Debug, Clone)]
182pub struct PromptTemplate {
183    source: String,
184    tokens: Vec<Token>,
185}
186
187impl PromptTemplate {
188    /// Parse a template string and validate variable names and filters.
189    ///
190    /// # Errors
191    ///
192    /// Returns [`TemplateError::ParseError`] if there are unclosed `{{` blocks,
193    /// or [`TemplateError::InvalidFilter`] for unrecognised filters.
194    #[allow(clippy::should_implement_trait)]
195    pub fn from_str(template: &str) -> Result<Self, TemplateError> {
196        let tokens = Self::parse(template)?;
197        Ok(Self {
198            source: template.to_string(),
199            tokens,
200        })
201    }
202
203    /// Return the original template source string.
204    pub fn source(&self) -> &str {
205        &self.source
206    }
207
208    /// Render the template using `ctx`.
209    ///
210    /// # Errors
211    ///
212    /// - [`TemplateError::MissingVariable`] — a variable has no binding and no
213    ///   `default` filter.
214    /// - [`TemplateError::InvalidFilter`] — an unrecognised filter was used.
215    pub fn render(&self, ctx: &TemplateContext) -> Result<String, TemplateError> {
216        let mut out = String::new();
217        for token in &self.tokens {
218            match token {
219                Token::Literal(s) => out.push_str(s),
220                Token::Substitution { variable, filters } => {
221                    let raw = ctx.get(variable);
222                    // Handle `default` filter when variable is missing.
223                    let base = match raw {
224                        Some(v) => v.to_string_repr(),
225                        None => {
226                            // Look for a default filter.
227                            let default_val = filters.iter().find_map(|f| {
228                                if let Filter::Default(d) = f {
229                                    Some(d.clone())
230                                } else {
231                                    None
232                                }
233                            });
234                            match default_val {
235                                Some(d) => d,
236                                None => {
237                                    return Err(TemplateError::MissingVariable(
238                                        variable.clone(),
239                                    ))
240                                }
241                            }
242                        }
243                    };
244                    // Apply non-default filters in order.
245                    let result = filters
246                        .iter()
247                        .filter(|f| !matches!(f, Filter::Default(_)))
248                        .fold(base, |acc, f| f.apply(&acc));
249                    out.push_str(&result);
250                }
251            }
252        }
253        Ok(out)
254    }
255
256    // ── Private helpers ───────────────────────────────────────────────────────
257
258    fn parse(template: &str) -> Result<Vec<Token>, TemplateError> {
259        let mut tokens = Vec::new();
260        let mut remaining = template;
261
262        while !remaining.is_empty() {
263            match remaining.find("{{") {
264                None => {
265                    // No more substitutions.
266                    tokens.push(Token::Literal(remaining.to_string()));
267                    break;
268                }
269                Some(start) => {
270                    if start > 0 {
271                        tokens.push(Token::Literal(remaining[..start].to_string()));
272                    }
273                    let after_open = &remaining[start + 2..];
274                    let end = after_open.find("}}").ok_or_else(|| {
275                        TemplateError::ParseError("unclosed '{{' block".to_string())
276                    })?;
277                    let inner = after_open[..end].trim();
278                    let token = Self::parse_substitution(inner)?;
279                    tokens.push(token);
280                    remaining = &after_open[end + 2..];
281                }
282            }
283        }
284
285        Ok(tokens)
286    }
287
288    fn parse_substitution(inner: &str) -> Result<Token, TemplateError> {
289        // Split on `|` to separate variable name from filters.
290        let parts: Vec<&str> = inner.splitn(2, '|').collect();
291        let variable = parts[0].trim().to_string();
292
293        // Validate variable name: must be a valid identifier.
294        if variable.is_empty()
295            || !variable
296                .chars()
297                .all(|c| c.is_alphanumeric() || c == '_')
298            || variable
299                .chars()
300                .next()
301                .map(|c| c.is_ascii_digit())
302                .unwrap_or(false)
303        {
304            return Err(TemplateError::ParseError(format!(
305                "invalid variable name: '{variable}'"
306            )));
307        }
308
309        let filters = if parts.len() > 1 {
310            // Multiple chained filters separated by `|`.
311            parts[1]
312                .split('|')
313                .map(|f| Filter::parse(f.trim()))
314                .collect::<Result<Vec<_>, _>>()?
315        } else {
316            Vec::new()
317        };
318
319        Ok(Token::Substitution { variable, filters })
320    }
321}
322
323// ── TemplateLibrary ───────────────────────────────────────────────────────────
324
325/// A named store of [`PromptTemplate`]s.
326///
327/// # Example
328///
329/// ```rust
330/// use tokio_prompt_orchestrator::template::{
331///     TemplateContext, TemplateLibrary, TemplateValue,
332/// };
333///
334/// let mut lib = TemplateLibrary::new();
335/// lib.register("greeting", "Hello, {{name}}!").unwrap();
336///
337/// let mut ctx = TemplateContext::new();
338/// ctx.set("name", TemplateValue::Text("Alice".to_string()));
339///
340/// let rendered = lib.render("greeting", &ctx).unwrap();
341/// assert_eq!(rendered, "Hello, Alice!");
342/// ```
343#[derive(Debug, Default)]
344pub struct TemplateLibrary {
345    templates: HashMap<String, PromptTemplate>,
346}
347
348impl TemplateLibrary {
349    /// Create an empty library.
350    pub fn new() -> Self {
351        Self::default()
352    }
353
354    /// Register a template under `name`.
355    ///
356    /// # Errors
357    ///
358    /// Returns a [`TemplateError`] if the template string cannot be parsed.
359    pub fn register(
360        &mut self,
361        name: impl Into<String>,
362        template: impl AsRef<str>,
363    ) -> Result<(), TemplateError> {
364        let t = PromptTemplate::from_str(template.as_ref())?;
365        self.templates.insert(name.into(), t);
366        Ok(())
367    }
368
369    /// Register a pre-compiled [`PromptTemplate`] under `name`.
370    pub fn register_template(&mut self, name: impl Into<String>, template: PromptTemplate) {
371        self.templates.insert(name.into(), template);
372    }
373
374    /// Render the named template with the given context.
375    ///
376    /// # Errors
377    ///
378    /// - [`TemplateError::MissingVariable`] if `name` is not registered.
379    /// - Any error from [`PromptTemplate::render`].
380    pub fn render(&self, name: &str, ctx: &TemplateContext) -> Result<String, TemplateError> {
381        let t = self
382            .templates
383            .get(name)
384            .ok_or_else(|| TemplateError::MissingVariable(name.to_string()))?;
385        t.render(ctx)
386    }
387
388    /// Return the names of all registered templates.
389    pub fn list(&self) -> Vec<&str> {
390        self.templates.keys().map(|s| s.as_str()).collect()
391    }
392
393    /// Return `true` if a template with the given name is registered.
394    pub fn contains(&self, name: &str) -> bool {
395        self.templates.contains_key(name)
396    }
397}
398
399// ── Tests ─────────────────────────────────────────────────────────────────────
400
401#[cfg(test)]
402#[allow(clippy::unwrap_used, clippy::expect_used)]
403mod tests {
404    use super::*;
405
406    fn ctx(pairs: &[(&str, &str)]) -> TemplateContext {
407        let mut c = TemplateContext::new();
408        for (k, v) in pairs {
409            c.set(*k, TemplateValue::Text(v.to_string()));
410        }
411        c
412    }
413
414    #[test]
415    fn test_plain_literal() {
416        let t = PromptTemplate::from_str("hello world").unwrap();
417        assert_eq!(t.render(&ctx(&[])).unwrap(), "hello world");
418    }
419
420    #[test]
421    fn test_simple_substitution() {
422        let t = PromptTemplate::from_str("Hello, {{name}}!").unwrap();
423        assert_eq!(t.render(&ctx(&[("name", "Alice")])).unwrap(), "Hello, Alice!");
424    }
425
426    #[test]
427    fn test_multiple_variables() {
428        let t = PromptTemplate::from_str("{{a}} + {{b}} = {{c}}").unwrap();
429        let c = ctx(&[("a", "1"), ("b", "2"), ("c", "3")]);
430        assert_eq!(t.render(&c).unwrap(), "1 + 2 = 3");
431    }
432
433    #[test]
434    fn test_missing_variable_error() {
435        let t = PromptTemplate::from_str("{{missing}}").unwrap();
436        let err = t.render(&ctx(&[])).unwrap_err();
437        assert_eq!(err, TemplateError::MissingVariable("missing".to_string()));
438    }
439
440    #[test]
441    fn test_filter_upper() {
442        let t = PromptTemplate::from_str("{{name | upper}}").unwrap();
443        assert_eq!(t.render(&ctx(&[("name", "hello")])).unwrap(), "HELLO");
444    }
445
446    #[test]
447    fn test_filter_lower() {
448        let t = PromptTemplate::from_str("{{name | lower}}").unwrap();
449        assert_eq!(t.render(&ctx(&[("name", "HELLO")])).unwrap(), "hello");
450    }
451
452    #[test]
453    fn test_filter_truncate() {
454        let t = PromptTemplate::from_str("{{text | truncate:5}}").unwrap();
455        assert_eq!(t.render(&ctx(&[("text", "hello world")])).unwrap(), "hello");
456    }
457
458    #[test]
459    fn test_filter_truncate_short_string() {
460        let t = PromptTemplate::from_str("{{text | truncate:100}}").unwrap();
461        assert_eq!(t.render(&ctx(&[("text", "hi")])).unwrap(), "hi");
462    }
463
464    #[test]
465    fn test_filter_default_present() {
466        let t = PromptTemplate::from_str(r#"{{name | default:"anon"}}"#).unwrap();
467        assert_eq!(t.render(&ctx(&[("name", "Bob")])).unwrap(), "Bob");
468    }
469
470    #[test]
471    fn test_filter_default_missing() {
472        let t = PromptTemplate::from_str(r#"{{name | default:"anon"}}"#).unwrap();
473        assert_eq!(t.render(&ctx(&[])).unwrap(), "anon");
474    }
475
476    #[test]
477    fn test_invalid_filter_name() {
478        let result = PromptTemplate::from_str("{{x | frobnicate}}");
479        assert!(matches!(result.unwrap_err(), TemplateError::InvalidFilter(_)));
480    }
481
482    #[test]
483    fn test_unclosed_brace_parse_error() {
484        let result = PromptTemplate::from_str("{{unclosed");
485        assert!(matches!(result.unwrap_err(), TemplateError::ParseError(_)));
486    }
487
488    #[test]
489    fn test_number_value() {
490        let t = PromptTemplate::from_str("Count: {{n}}").unwrap();
491        let mut c = TemplateContext::new();
492        c.set("n", TemplateValue::Number(42.0));
493        assert_eq!(t.render(&c).unwrap(), "Count: 42");
494    }
495
496    #[test]
497    fn test_bool_value() {
498        let t = PromptTemplate::from_str("Active: {{flag}}").unwrap();
499        let mut c = TemplateContext::new();
500        c.set("flag", TemplateValue::Bool(true));
501        assert_eq!(t.render(&c).unwrap(), "Active: true");
502    }
503
504    #[test]
505    fn test_list_value() {
506        let t = PromptTemplate::from_str("Items: {{list}}").unwrap();
507        let mut c = TemplateContext::new();
508        c.set(
509            "list",
510            TemplateValue::List(vec!["a".to_string(), "b".to_string(), "c".to_string()]),
511        );
512        assert_eq!(t.render(&c).unwrap(), "Items: a, b, c");
513    }
514
515    #[test]
516    fn test_library_register_and_render() {
517        let mut lib = TemplateLibrary::new();
518        lib.register("greet", "Hi {{name}}!").unwrap();
519        let rendered = lib.render("greet", &ctx(&[("name", "World")])).unwrap();
520        assert_eq!(rendered, "Hi World!");
521    }
522
523    #[test]
524    fn test_library_list() {
525        let mut lib = TemplateLibrary::new();
526        lib.register("a", "{{x}}").unwrap();
527        lib.register("b", "{{y}}").unwrap();
528        let mut names = lib.list();
529        names.sort();
530        assert_eq!(names, vec!["a", "b"]);
531    }
532
533    #[test]
534    fn test_library_missing_template() {
535        let lib = TemplateLibrary::new();
536        let err = lib.render("nonexistent", &ctx(&[])).unwrap_err();
537        assert!(matches!(err, TemplateError::MissingVariable(_)));
538    }
539
540    #[test]
541    fn test_library_contains() {
542        let mut lib = TemplateLibrary::new();
543        lib.register("t", "x").unwrap();
544        assert!(lib.contains("t"));
545        assert!(!lib.contains("missing"));
546    }
547
548    #[test]
549    fn test_chained_filters_upper_truncate() {
550        let t = PromptTemplate::from_str("{{text | upper | truncate:3}}").unwrap();
551        assert_eq!(t.render(&ctx(&[("text", "hello")])).unwrap(), "HEL");
552    }
553
554    #[test]
555    fn test_no_substitutions() {
556        let t = PromptTemplate::from_str("plain text").unwrap();
557        assert_eq!(t.render(&ctx(&[])).unwrap(), "plain text");
558    }
559
560    #[test]
561    fn test_adjacent_substitutions() {
562        let t = PromptTemplate::from_str("{{a}}{{b}}").unwrap();
563        assert_eq!(t.render(&ctx(&[("a", "foo"), ("b", "bar")])).unwrap(), "foobar");
564    }
565
566    #[test]
567    fn test_empty_template() {
568        let t = PromptTemplate::from_str("").unwrap();
569        assert_eq!(t.render(&ctx(&[])).unwrap(), "");
570    }
571
572    #[test]
573    fn test_source_preserved() {
574        let src = "Hello, {{name}}!";
575        let t = PromptTemplate::from_str(src).unwrap();
576        assert_eq!(t.source(), src);
577    }
578}