1use std::collections::HashMap;
7use std::fmt;
8
9#[derive(Debug, Clone, PartialEq)]
13pub enum TemplateError {
14 UnclosedBlock,
16 UnknownVariable(String),
18 InvalidSyntax(String),
20 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#[derive(Debug, Clone, PartialEq)]
41pub enum TemplateVar {
42 Str(String),
44 Int(i64),
46 Float(f64),
48 Bool(bool),
50 List(Vec<TemplateVar>),
52 Map(HashMap<String, TemplateVar>),
54}
55
56impl TemplateVar {
57 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 pub fn as_bool(&self) -> bool {
77 self.is_truthy()
78 }
79
80 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#[derive(Debug, Clone, Default)]
97pub struct TemplateContext {
98 vars: HashMap<String, TemplateVar>,
99}
100
101impl TemplateContext {
102 pub fn new() -> Self {
104 Self::default()
105 }
106
107 pub fn set(&mut self, key: &str, val: TemplateVar) {
109 self.vars.insert(key.to_string(), val);
110 }
111
112 pub fn get(&self, key: &str) -> Option<&TemplateVar> {
114 self.vars.get(key)
115 }
116
117 pub fn set_str(&mut self, key: &str, val: &str) {
119 self.set(key, TemplateVar::Str(val.to_string()));
120 }
121
122 pub fn set_int(&mut self, key: &str, val: i64) {
124 self.set(key, TemplateVar::Int(val));
125 }
126
127 pub fn set_bool(&mut self, key: &str, val: bool) {
129 self.set(key, TemplateVar::Bool(val));
130 }
131
132 pub fn set_list(&mut self, key: &str, val: Vec<TemplateVar>) {
134 self.set(key, TemplateVar::List(val));
135 }
136}
137
138#[derive(Debug, Clone)]
142enum Node {
143 Text(String),
145 Expr(String),
147 If {
149 branches: Vec<(String, Vec<Node>)>, else_branch: Option<Vec<Node>>,
151 },
152 For {
154 item: String,
155 list: String,
156 body: Vec<Node>,
157 },
158}
159
160#[derive(Debug, Clone)]
164pub struct Template {
165 source: String,
166 nodes: Vec<Node>,
167}
168
169const MAX_NEST: usize = 32;
170
171impl Template {
172 pub fn new(source: &str) -> Result<Self, TemplateError> {
174 let nodes = parse(source)?;
175 Ok(Self { source: source.to_string(), nodes })
176 }
177
178 pub fn render(&self, ctx: &TemplateContext) -> Result<String, TemplateError> {
180 render_nodes(&self.nodes, ctx)
181 }
182
183 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 pub fn validate(&self) -> Result<(), TemplateError> {
194 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 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#[derive(Debug, Clone)]
235enum Token {
236 Text(String),
237 Expr(String), Tag(String), Comment, }
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 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
275fn 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 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 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 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
388fn 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
435fn eval_expr(expr: &str, ctx: &TemplateContext) -> Result<String, TemplateError> {
437 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 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 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 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 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 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 if let Some(inner) = cond.strip_prefix("not ") {
522 return eval_condition(inner.trim(), ctx).map(|b| !b);
523 }
524
525 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 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 match ctx.get(cond) {
543 Some(v) => Ok(v.is_truthy()),
544 None => Ok(false),
545 }
546}
547
548#[derive(Debug, Default)]
552pub struct TemplateLibrary {
553 templates: HashMap<String, Template>,
554}
555
556impl TemplateLibrary {
557 pub fn new() -> Self {
559 Self::default()
560 }
561
562 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 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 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}