tokio_prompt_orchestrator/
template.rs1use std::collections::HashMap;
28use thiserror::Error;
29
30#[derive(Debug, Error, PartialEq, Clone)]
34pub enum TemplateError {
35 #[error("missing variable: {0}")]
38 MissingVariable(String),
39
40 #[error("invalid filter: {0}")]
42 InvalidFilter(String),
43
44 #[error("parse error: {0}")]
46 ParseError(String),
47}
48
49#[derive(Debug, Clone, PartialEq)]
53pub enum TemplateValue {
54 Text(String),
56 Number(f64),
58 Bool(bool),
60 List(Vec<String>),
62}
63
64impl TemplateValue {
65 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#[derive(Debug, Clone, Default)]
86pub struct TemplateContext {
87 values: HashMap<String, TemplateValue>,
88}
89
90impl TemplateContext {
91 pub fn new() -> Self {
93 Self::default()
94 }
95
96 pub fn set(&mut self, key: impl Into<String>, value: TemplateValue) {
98 self.values.insert(key.into(), value);
99 }
100
101 pub fn get(&self, key: &str) -> Option<&TemplateValue> {
103 self.values.get(key)
104 }
105}
106
107#[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 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(), }
161 }
162}
163
164#[derive(Debug, Clone)]
167enum Token {
168 Literal(String),
169 Substitution {
170 variable: String,
171 filters: Vec<Filter>,
172 },
173}
174
175#[derive(Debug, Clone)]
182pub struct PromptTemplate {
183 source: String,
184 tokens: Vec<Token>,
185}
186
187impl PromptTemplate {
188 #[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 pub fn source(&self) -> &str {
205 &self.source
206 }
207
208 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 let base = match raw {
224 Some(v) => v.to_string_repr(),
225 None => {
226 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 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 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 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 let parts: Vec<&str> = inner.splitn(2, '|').collect();
291 let variable = parts[0].trim().to_string();
292
293 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 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#[derive(Debug, Default)]
344pub struct TemplateLibrary {
345 templates: HashMap<String, PromptTemplate>,
346}
347
348impl TemplateLibrary {
349 pub fn new() -> Self {
351 Self::default()
352 }
353
354 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 pub fn register_template(&mut self, name: impl Into<String>, template: PromptTemplate) {
371 self.templates.insert(name.into(), template);
372 }
373
374 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 pub fn list(&self) -> Vec<&str> {
390 self.templates.keys().map(|s| s.as_str()).collect()
391 }
392
393 pub fn contains(&self, name: &str) -> bool {
395 self.templates.contains_key(name)
396 }
397}
398
399#[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}