Skip to main content

tokio_prompt_orchestrator/
templates.rs

1//! # Prompt Template Engine
2//!
3//! Versioned, variable-interpolated prompt templates with built-in A/B testing.
4//!
5//! ## Overview
6//!
7//! - **Templates** — named prompt blueprints with `{{variable}}` placeholders,
8//!   optional system prompts, and version tags.
9//! - **Registry** — hot-reloadable store for all templates; load from TOML.
10//! - **A/B Experiments** — route traffic between template variants, collect
11//!   per-variant latency and quality metrics, and determine winners via a
12//!   two-proportion Z-test.
13//!
14//! ## Quick start
15//!
16//! ```rust,no_run
17//! use tokio_prompt_orchestrator::templates::{PromptTemplate, TemplateRegistry};
18//! use std::collections::HashMap;
19//!
20//! let registry = TemplateRegistry::new();
21//!
22//! let t = PromptTemplate::builder("summarise")
23//!     .version("v2")
24//!     .system("You are a concise summariser.")
25//!     .body("Summarise the following in {{max_words}} words or fewer:\n\n{{text}}")
26//!     .tag("summarisation")
27//!     .build();
28//!
29//! registry.register(t);
30//!
31//! let mut vars = HashMap::new();
32//! vars.insert("max_words", "50");
33//! vars.insert("text", "The quick brown fox…");
34//!
35//! let rendered = registry.render("summarise", &vars).unwrap();
36//! println!("{rendered}");
37//! ```
38
39use std::{
40    collections::HashMap,
41    sync::{Arc, RwLock},
42};
43
44// ---------------------------------------------------------------------------
45// Template
46// ---------------------------------------------------------------------------
47
48/// A versioned prompt template with variable placeholders.
49///
50/// Placeholders use `{{variable_name}}` syntax. Variables are resolved at
51/// render time from a caller-supplied map. Missing variables default to the
52/// empty string unless a default is configured on the template.
53#[derive(Debug, Clone)]
54pub struct PromptTemplate {
55    /// Unique template name used as the lookup key.
56    pub name: String,
57    /// Semantic version string (e.g. `"v1"`, `"v2.1"`).
58    pub version: String,
59    /// Optional system prompt (sent as the first turn or prefixed).
60    pub system: Option<String>,
61    /// Template body with `{{variable}}` placeholders.
62    pub body: String,
63    /// Variable names declared in this template with optional defaults.
64    pub variables: HashMap<String, Option<String>>,
65    /// Arbitrary tags for grouping and filtering.
66    pub tags: Vec<String>,
67    /// Human-readable description.
68    pub description: Option<String>,
69}
70
71impl PromptTemplate {
72    /// Start building a template with the given name.
73    pub fn builder(name: impl Into<String>) -> TemplateBuilder {
74        TemplateBuilder::new(name)
75    }
76
77    /// Render the template body by substituting `vars` for each `{{key}}`.
78    ///
79    /// In addition to simple `{{variable}}` interpolation this method handles:
80    ///
81    /// - **Conditional blocks** `{{#if condition}}...{{/if}}` — the block is
82    ///   included if `vars` contains a key equal to `condition` whose value is
83    ///   non-empty and not the literal string `"false"` or `"0"`.
84    ///
85    /// - **Loop blocks** `{{#each items}}...{{/each}}` — the block is repeated
86    ///   once for each item in a comma-separated list stored under the key
87    ///   `items`.  Inside the block `{{this}}` is replaced with the current
88    ///   item and `{{@index}}` with the zero-based index.
89    ///
90    /// Variables present in the template but absent from `vars` fall back to
91    /// the template's own default value. If no default exists the placeholder
92    /// is replaced with an empty string.
93    pub fn render(&self, vars: &HashMap<&str, &str>) -> String {
94        // Step 1 — process {{#each items}}...{{/each}} blocks
95        let mut output = render_each_blocks(&self.body, vars);
96
97        // Step 2 — process {{#if condition}}...{{/if}} blocks
98        output = render_if_blocks(&output, vars);
99
100        // Step 3 — simple variable substitution
101        for (key, default) in &self.variables {
102            let placeholder = format!("{{{{{key}}}}}");
103            let value = vars
104                .get(key.as_str())
105                .copied()
106                .or(default.as_deref())
107                .unwrap_or("");
108            output = output.replace(&placeholder, value);
109        }
110        // Also substitute any vars not pre-declared
111        for (key, value) in vars {
112            let placeholder = format!("{{{{{key}}}}}");
113            if output.contains(&placeholder) {
114                output = output.replace(&placeholder, value);
115            }
116        }
117        output
118    }
119
120    /// Render both system prompt (if any) and body; returns `(system, body)`.
121    pub fn render_full(
122        &self,
123        vars: &HashMap<&str, &str>,
124    ) -> (Option<String>, String) {
125        let system = self.system.as_ref().map(|s| {
126            let mut out = s.clone();
127            for (k, v) in vars {
128                out = out.replace(&format!("{{{{{k}}}}}"), v);
129            }
130            out
131        });
132        (system, self.render(vars))
133    }
134
135    /// Extract variable names declared in the body (`{{name}}`).
136    pub fn declared_variables(&self) -> Vec<String> {
137        extract_placeholders(&self.body)
138    }
139}
140
141// ---------------------------------------------------------------------------
142// Builder
143// ---------------------------------------------------------------------------
144
145/// Fluent builder for [`PromptTemplate`].
146#[derive(Default)]
147pub struct TemplateBuilder {
148    name: String,
149    version: String,
150    system: Option<String>,
151    body: String,
152    defaults: HashMap<String, Option<String>>,
153    tags: Vec<String>,
154    description: Option<String>,
155}
156
157impl TemplateBuilder {
158    fn new(name: impl Into<String>) -> Self {
159        TemplateBuilder {
160            name: name.into(),
161            version: "v1".into(),
162            ..Default::default()
163        }
164    }
165
166    pub fn version(mut self, v: impl Into<String>) -> Self {
167        self.version = v.into();
168        self
169    }
170
171    pub fn system(mut self, s: impl Into<String>) -> Self {
172        self.system = Some(s.into());
173        self
174    }
175
176    pub fn body(mut self, b: impl Into<String>) -> Self {
177        self.body = b.into();
178        self
179    }
180
181    /// Declare a required variable (no default).
182    pub fn var(mut self, name: impl Into<String>) -> Self {
183        self.defaults.insert(name.into(), None);
184        self
185    }
186
187    /// Declare a variable with a fallback default value.
188    pub fn var_default(mut self, name: impl Into<String>, default: impl Into<String>) -> Self {
189        self.defaults.insert(name.into(), Some(default.into()));
190        self
191    }
192
193    pub fn tag(mut self, t: impl Into<String>) -> Self {
194        self.tags.push(t.into());
195        self
196    }
197
198    pub fn description(mut self, d: impl Into<String>) -> Self {
199        self.description = Some(d.into());
200        self
201    }
202
203    /// Finalise and return the [`PromptTemplate`].
204    ///
205    /// Auto-detects variables from `{{…}}` placeholders in the body if they
206    /// haven't been explicitly declared.
207    pub fn build(mut self) -> PromptTemplate {
208        for v in extract_placeholders(&self.body) {
209            self.defaults.entry(v).or_insert(None);
210        }
211        if let Some(sys) = &self.system {
212            for v in extract_placeholders(sys) {
213                self.defaults.entry(v).or_insert(None);
214            }
215        }
216        PromptTemplate {
217            name: self.name,
218            version: self.version,
219            system: self.system,
220            body: self.body,
221            variables: self.defaults,
222            tags: self.tags,
223            description: self.description,
224        }
225    }
226}
227
228// ---------------------------------------------------------------------------
229// Registry
230// ---------------------------------------------------------------------------
231
232/// Template lookup errors.
233#[derive(Debug, thiserror::Error)]
234pub enum TemplateError {
235    #[error("template '{0}' not found")]
236    NotFound(String),
237    #[error("missing required variable '{0}'")]
238    MissingVariable(String),
239}
240
241/// Thread-safe registry of named prompt templates.
242///
243/// Cheap to clone (Arc-backed). The registry supports hot-reload: call
244/// [`register`](Self::register) to add or replace a template at any time.
245#[derive(Clone, Default)]
246pub struct TemplateRegistry {
247    inner: Arc<RwLock<HashMap<String, PromptTemplate>>>,
248}
249
250impl TemplateRegistry {
251    pub fn new() -> Self {
252        TemplateRegistry::default()
253    }
254
255    /// Register (or replace) a template.
256    pub fn register(&self, template: PromptTemplate) {
257        let mut guard = self.inner.write().unwrap_or_else(|e| e.into_inner());
258        guard.insert(template.name.clone(), template);
259    }
260
261    /// Remove a template by name.
262    pub fn unregister(&self, name: &str) {
263        let mut guard = self.inner.write().unwrap_or_else(|e| e.into_inner());
264        guard.remove(name);
265    }
266
267    /// Render a named template with the supplied variables.
268    pub fn render(
269        &self,
270        name: &str,
271        vars: &HashMap<&str, &str>,
272    ) -> Result<String, TemplateError> {
273        let guard = self.inner.read().unwrap_or_else(|e| e.into_inner());
274        let tmpl = guard.get(name).ok_or_else(|| TemplateError::NotFound(name.to_owned()))?;
275
276        // Validate required variables
277        for (var, default) in &tmpl.variables {
278            if default.is_none() && !vars.contains_key(var.as_str()) {
279                return Err(TemplateError::MissingVariable(var.clone()));
280            }
281        }
282        Ok(tmpl.render(vars))
283    }
284
285    /// Render a template and return `(system_prompt, body)`.
286    pub fn render_full(
287        &self,
288        name: &str,
289        vars: &HashMap<&str, &str>,
290    ) -> Result<(Option<String>, String), TemplateError> {
291        let guard = self.inner.read().unwrap_or_else(|e| e.into_inner());
292        let tmpl = guard.get(name).ok_or_else(|| TemplateError::NotFound(name.to_owned()))?;
293        Ok(tmpl.render_full(vars))
294    }
295
296    /// Get a snapshot of a template (cloned).
297    pub fn get(&self, name: &str) -> Option<PromptTemplate> {
298        let guard = self.inner.read().unwrap_or_else(|e| e.into_inner());
299        guard.get(name).cloned()
300    }
301
302    /// List all registered template names.
303    pub fn names(&self) -> Vec<String> {
304        let guard = self.inner.read().unwrap_or_else(|e| e.into_inner());
305        guard.keys().cloned().collect()
306    }
307
308    /// Load templates from a TOML string.
309    ///
310    /// Expected format:
311    /// ```toml
312    /// [[templates]]
313    /// name    = "summarise"
314    /// version = "v1"
315    /// system  = "You are a summariser."
316    /// body    = "Summarise: {{text}}"
317    /// tags    = ["nlp"]
318    ///
319    /// [[templates]]
320    /// name = "translate"
321    /// body = "Translate to {{lang}}: {{text}}"
322    /// ```
323    pub fn load_toml(&self, toml_str: &str) -> Result<usize, String> {
324        let parsed: toml::Value =
325            toml::from_str(toml_str).map_err(|e| format!("TOML parse error: {e}"))?;
326
327        let arr = parsed
328            .get("templates")
329            .and_then(|v| v.as_array())
330            .ok_or("expected [[templates]] array")?;
331
332        let mut count = 0;
333        for item in arr {
334            let name = item
335                .get("name")
336                .and_then(|v| v.as_str())
337                .ok_or("template missing 'name'")?
338                .to_owned();
339
340            let version = item
341                .get("version")
342                .and_then(|v| v.as_str())
343                .unwrap_or("v1")
344                .to_owned();
345
346            let body = item
347                .get("body")
348                .and_then(|v| v.as_str())
349                .ok_or(format!("template '{name}' missing 'body'"))?
350                .to_owned();
351
352            let system = item.get("system").and_then(|v| v.as_str()).map(str::to_owned);
353
354            let description = item
355                .get("description")
356                .and_then(|v| v.as_str())
357                .map(str::to_owned);
358
359            let tags: Vec<String> = item
360                .get("tags")
361                .and_then(|v| v.as_array())
362                .map(|arr| {
363                    arr.iter()
364                        .filter_map(|v| v.as_str())
365                        .map(str::to_owned)
366                        .collect()
367                })
368                .unwrap_or_default();
369
370            let mut builder = PromptTemplate::builder(name)
371                .version(version)
372                .body(body);
373
374            if let Some(sys) = system {
375                builder = builder.system(sys);
376            }
377            if let Some(desc) = description {
378                builder = builder.description(desc);
379            }
380            for tag in tags {
381                builder = builder.tag(tag);
382            }
383
384            self.register(builder.build());
385            count += 1;
386        }
387        Ok(count)
388    }
389}
390
391// ---------------------------------------------------------------------------
392// A/B Testing
393// ---------------------------------------------------------------------------
394
395/// A single variant in an A/B experiment.
396#[derive(Debug, Clone)]
397pub struct ExperimentVariant {
398    /// Template name in the registry.
399    pub template_name: String,
400    /// Traffic weight relative to other variants (unnormalized).
401    pub weight: f64,
402    /// Display label for this variant.
403    pub label: String,
404}
405
406/// Metrics collected per A/B variant.
407#[derive(Debug, Clone, Default)]
408pub struct VariantMetrics {
409    pub requests: u64,
410    pub successes: u64,
411    pub total_latency_ms: u64,
412    /// Accumulated quality score (caller-reported, 0.0–1.0 per request).
413    pub total_quality: f64,
414}
415
416impl VariantMetrics {
417    pub fn success_rate(&self) -> f64 {
418        if self.requests == 0 {
419            return 0.0;
420        }
421        self.successes as f64 / self.requests as f64
422    }
423
424    pub fn avg_latency_ms(&self) -> f64 {
425        if self.requests == 0 {
426            return 0.0;
427        }
428        self.total_latency_ms as f64 / self.requests as f64
429    }
430
431    pub fn avg_quality(&self) -> f64 {
432        if self.successes == 0 {
433            return 0.0;
434        }
435        self.total_quality / self.successes as f64
436    }
437}
438
439/// An A/B experiment that routes traffic between template variants.
440pub struct AbExperiment {
441    pub name: String,
442    pub variants: Vec<ExperimentVariant>,
443    metrics: Arc<RwLock<Vec<VariantMetrics>>>,
444    cumulative_weights: Vec<f64>,
445    total_weight: f64,
446}
447
448impl AbExperiment {
449    /// Create a new A/B experiment.
450    ///
451    /// `variants` must not be empty.
452    ///
453    /// # Panics
454    ///
455    /// Panics if `variants` is empty.
456    pub fn new(name: impl Into<String>, variants: Vec<ExperimentVariant>) -> Self {
457        assert!(!variants.is_empty(), "experiment must have at least one variant");
458        let n = variants.len();
459        let mut cum = 0.0f64;
460        let mut cumulative_weights = Vec::with_capacity(n);
461        for v in &variants {
462            cum += v.weight;
463            cumulative_weights.push(cum);
464        }
465        let metrics = vec![VariantMetrics::default(); n];
466        AbExperiment {
467            name: name.into(),
468            cumulative_weights,
469            total_weight: cum,
470            variants,
471            metrics: Arc::new(RwLock::new(metrics)),
472        }
473    }
474
475    /// Select a variant index using a [0.0, 1.0) uniform random draw.
476    ///
477    /// Pass a caller-supplied random value so callers can use any PRNG.
478    pub fn pick_variant(&self, rand_0_1: f64) -> usize {
479        let target = rand_0_1 * self.total_weight;
480        for (i, &w) in self.cumulative_weights.iter().enumerate() {
481            if target < w {
482                return i;
483            }
484        }
485        self.variants.len() - 1
486    }
487
488    /// Record that a request was routed to variant `idx`.
489    pub fn record_request(&self, idx: usize) {
490        let mut guard = self.metrics.write().unwrap_or_else(|e| e.into_inner());
491        if let Some(m) = guard.get_mut(idx) {
492            m.requests += 1;
493        }
494    }
495
496    /// Record a successful response for variant `idx`.
497    pub fn record_success(&self, idx: usize, latency_ms: u64, quality: f64) {
498        let mut guard = self.metrics.write().unwrap_or_else(|e| e.into_inner());
499        if let Some(m) = guard.get_mut(idx) {
500            m.successes += 1;
501            m.total_latency_ms += latency_ms;
502            m.total_quality += quality.clamp(0.0, 1.0);
503        }
504    }
505
506    /// Return a snapshot of all variant metrics.
507    pub fn metrics(&self) -> Vec<VariantMetrics> {
508        self.metrics.read().unwrap_or_else(|e| e.into_inner()).clone()
509    }
510
511    /// Return the index of the leading variant by quality score.
512    /// Returns `None` if no successes have been recorded.
513    pub fn leading_variant(&self) -> Option<usize> {
514        let guard = self.metrics.read().unwrap_or_else(|e| e.into_inner());
515        guard
516            .iter()
517            .enumerate()
518            .filter(|(_, m)| m.successes > 0)
519            .max_by(|(_, a), (_, b)| a.avg_quality().total_cmp(&b.avg_quality()))
520            .map(|(i, _)| i)
521    }
522
523    /// Two-proportion Z-test p-value between variants `a` and `b` (success rate).
524    ///
525    /// Returns `None` if sample sizes are insufficient (< 30 per variant).
526    pub fn significance(&self, a: usize, b: usize) -> Option<f64> {
527        let guard = self.metrics.read().unwrap_or_else(|e| e.into_inner());
528        let ma = guard.get(a)?;
529        let mb = guard.get(b)?;
530        if ma.requests < 30 || mb.requests < 30 {
531            return None;
532        }
533        let pa = ma.success_rate();
534        let pb = mb.success_rate();
535        let na = ma.requests as f64;
536        let nb = mb.requests as f64;
537        let p_pool = (ma.successes + mb.successes) as f64 / (na + nb);
538        let denom = (p_pool * (1.0 - p_pool) * (1.0 / na + 1.0 / nb)).sqrt();
539        if denom < f64::EPSILON {
540            return None;
541        }
542        let z = (pa - pb).abs() / denom;
543        // Approximate two-tailed p-value via standard normal CDF
544        let p = 2.0 * standard_normal_cdf(-z.abs());
545        Some(p)
546    }
547
548    /// Summary report for logging / monitoring.
549    pub fn report(&self) -> ExperimentReport {
550        let metrics = self.metrics();
551        let variants = self
552            .variants
553            .iter()
554            .zip(metrics.iter())
555            .map(|(v, m)| VariantReport {
556                label: v.label.clone(),
557                template_name: v.template_name.clone(),
558                requests: m.requests,
559                success_rate: m.success_rate(),
560                avg_latency_ms: m.avg_latency_ms(),
561                avg_quality: m.avg_quality(),
562            })
563            .collect();
564        ExperimentReport {
565            experiment: self.name.clone(),
566            variants,
567            leading_variant: self.leading_variant().map(|i| self.variants[i].label.clone()),
568        }
569    }
570}
571
572/// A point-in-time report for an A/B experiment.
573#[derive(Debug, Clone, serde::Serialize)]
574pub struct ExperimentReport {
575    pub experiment: String,
576    pub variants: Vec<VariantReport>,
577    pub leading_variant: Option<String>,
578}
579
580#[derive(Debug, Clone, serde::Serialize)]
581pub struct VariantReport {
582    pub label: String,
583    pub template_name: String,
584    pub requests: u64,
585    pub success_rate: f64,
586    pub avg_latency_ms: f64,
587    pub avg_quality: f64,
588}
589
590// ---------------------------------------------------------------------------
591// Helpers
592// ---------------------------------------------------------------------------
593
594// ---------------------------------------------------------------------------
595// Block rendering helpers
596// ---------------------------------------------------------------------------
597
598/// Evaluate a condition value: truthy when non-empty and not `"false"`/`"0"`.
599fn is_truthy(value: &str) -> bool {
600    !value.is_empty() && value != "false" && value != "0"
601}
602
603/// Process all `{{#if condition}}...{{/if}}` blocks in `template`.
604///
605/// - The block is **included** (without the tags) when the condition is truthy.
606/// - The block is **removed** (including the tags) when the condition is falsy.
607/// - Nested `{{#if}}` blocks are **not** supported — they are left as-is.
608fn render_if_blocks(template: &str, vars: &HashMap<&str, &str>) -> String {
609    let mut output = template.to_string();
610
611    while let Some(open_start) = output.find("{{#if ") {
612        // Found the next {{#if <condition>}} tag
613        let tag_end = match output[open_start..].find("}}") {
614            Some(j) => open_start + j + 2,
615            None => break,
616        };
617
618        // Extract the condition name from between "{{#if " and "}}"
619        let condition = output[open_start + 6..tag_end - 2].trim().to_string();
620
621        // Locate the matching {{/if}}
622        let close_tag = "{{/if}}";
623        let close_start = match output[tag_end..].find(close_tag) {
624            Some(k) => tag_end + k,
625            None => break, // malformed template — stop processing
626        };
627        let close_end = close_start + close_tag.len();
628
629        // Body is the content between the opening and closing tags
630        let body = output[tag_end..close_start].to_string();
631
632        // Decide what to replace the whole block with
633        let condition_value = vars.get(condition.as_str()).copied().unwrap_or("");
634        let replacement = if is_truthy(condition_value) {
635            body
636        } else {
637            String::new()
638        };
639
640        output = format!("{}{}{}", &output[..open_start], replacement, &output[close_end..]);
641    }
642
643    output
644}
645
646/// Process all `{{#each items}}...{{/each}}` blocks in `template`.
647///
648/// The value of `items` in `vars` is treated as a **comma-separated list**.
649/// Each element is trimmed before use.  Inside the block:
650///
651/// - `{{this}}` — replaced with the current item's value.
652/// - `{{@index}}` — replaced with the zero-based integer index.
653///
654/// If `items` is absent or empty the entire block is removed.
655fn render_each_blocks(template: &str, vars: &HashMap<&str, &str>) -> String {
656    let mut output = template.to_string();
657
658    while let Some(open_start) = output.find("{{#each ") {
659        // Found the next {{#each <list_key>}} tag
660        let tag_end = match output[open_start..].find("}}") {
661            Some(j) => open_start + j + 2,
662            None => break,
663        };
664
665        // Extract the list key from between "{{#each " and "}}"
666        let list_key = output[open_start + 8..tag_end - 2].trim().to_string();
667
668        // Locate the matching {{/each}}
669        let close_tag = "{{/each}}";
670        let close_start = match output[tag_end..].find(close_tag) {
671            Some(k) => tag_end + k,
672            None => break, // malformed — stop
673        };
674        let close_end = close_start + close_tag.len();
675
676        // The body template for each iteration
677        let body_template = output[tag_end..close_start].to_string();
678
679        // Expand
680        let items_str = vars.get(list_key.as_str()).copied().unwrap_or("");
681        let mut expanded = String::new();
682
683        if !items_str.is_empty() {
684            for (idx, item) in items_str.split(',').enumerate() {
685                let item = item.trim();
686                let iter_body = body_template
687                    .replace("{{this}}", item)
688                    .replace("{{@index}}", &idx.to_string());
689                expanded.push_str(&iter_body);
690            }
691        }
692
693        output = format!("{}{}{}", &output[..open_start], expanded, &output[close_end..]);
694    }
695
696    output
697}
698
699/// Extract placeholder names from a `{{…}}` template string.
700fn extract_placeholders(s: &str) -> Vec<String> {
701    let mut vars = Vec::new();
702    let mut chars = s.char_indices().peekable();
703    while let Some((i, c)) = chars.next() {
704        if c == '{' {
705            if let Some((_, '{')) = chars.peek() {
706                chars.next();
707                let start = i + 2;
708                let mut end = start;
709                while let Some(&(j, c2)) = chars.peek() {
710                    if c2 == '}' {
711                        chars.next();
712                        if let Some(&(_, '}')) = chars.peek() {
713                            chars.next();
714                            end = j;
715                            break;
716                        }
717                    } else {
718                        end = j + c2.len_utf8();
719                        chars.next();
720                    }
721                }
722                let name = s[start..end].trim().to_owned();
723                if !name.is_empty() && !vars.contains(&name) {
724                    vars.push(name);
725                }
726            }
727        }
728    }
729    vars
730}
731
732/// Approximation of the standard normal CDF Φ(x).
733/// Uses the Abramowitz & Stegun rational approximation (max error ≈ 7.5×10⁻⁸).
734fn standard_normal_cdf(x: f64) -> f64 {
735    let t = 1.0 / (1.0 + 0.2316419 * x.abs());
736    let poly = t
737        * (0.319381530
738            + t * (-0.356563782
739                + t * (1.781477937
740                    + t * (-1.821255978 + t * 1.330274429))));
741    let phi = 1.0 - (-(x * x / 2.0)).exp() / (2.0 * std::f64::consts::PI).sqrt() * poly;
742    if x >= 0.0 {
743        phi
744    } else {
745        1.0 - phi
746    }
747}
748
749// ---------------------------------------------------------------------------
750// Tests
751// ---------------------------------------------------------------------------
752
753#[cfg(test)]
754mod tests {
755    use super::*;
756
757    #[test]
758    fn test_basic_render() {
759        let t = PromptTemplate::builder("t1")
760            .body("Hello {{name}}, you are {{age}} years old.")
761            .build();
762        let mut vars = HashMap::new();
763        vars.insert("name", "Alice");
764        vars.insert("age", "30");
765        assert_eq!(t.render(&vars), "Hello Alice, you are 30 years old.");
766    }
767
768    #[test]
769    fn test_default_variable() {
770        let t = PromptTemplate::builder("t2")
771            .body("Format: {{style}}")
772            .var_default("style", "markdown")
773            .build();
774        let vars = HashMap::new();
775        assert_eq!(t.render(&vars), "Format: markdown");
776    }
777
778    #[test]
779    fn test_missing_required_variable_error() {
780        let registry = TemplateRegistry::new();
781        let t = PromptTemplate::builder("t3")
782            .body("Hello {{name}}")
783            .var("name")
784            .build();
785        registry.register(t);
786        let vars: HashMap<&str, &str> = HashMap::new();
787        assert!(matches!(
788            registry.render("t3", &vars),
789            Err(TemplateError::MissingVariable(_))
790        ));
791    }
792
793    #[test]
794    fn test_not_found_error() {
795        let registry = TemplateRegistry::new();
796        let vars: HashMap<&str, &str> = HashMap::new();
797        assert!(matches!(
798            registry.render("nonexistent", &vars),
799            Err(TemplateError::NotFound(_))
800        ));
801    }
802
803    #[test]
804    fn test_extract_placeholders() {
805        let vars = extract_placeholders("Say {{greeting}} to {{name}}!");
806        assert_eq!(vars, vec!["greeting", "name"]);
807    }
808
809    #[test]
810    fn test_toml_load() {
811        let registry = TemplateRegistry::new();
812        let toml = r#"
813[[templates]]
814name    = "greet"
815version = "v1"
816body    = "Hello {{name}}!"
817tags    = ["greeting"]
818"#;
819        let count = registry.load_toml(toml).unwrap();
820        assert_eq!(count, 1);
821        let mut vars = HashMap::new();
822        vars.insert("name", "Bob");
823        assert_eq!(registry.render("greet", &vars).unwrap(), "Hello Bob!");
824    }
825
826    #[test]
827    fn test_ab_experiment_routing() {
828        let variants = vec![
829            ExperimentVariant {
830                template_name: "tmpl_a".into(),
831                weight: 70.0,
832                label: "control".into(),
833            },
834            ExperimentVariant {
835                template_name: "tmpl_b".into(),
836                weight: 30.0,
837                label: "treatment".into(),
838            },
839        ];
840        let exp = AbExperiment::new("test_exp", variants);
841
842        // 0.0 should always pick first variant (weight 70/100)
843        assert_eq!(exp.pick_variant(0.0), 0);
844        // 0.99 should pick second variant
845        assert_eq!(exp.pick_variant(0.99), 1);
846    }
847
848    #[test]
849    fn test_ab_metrics() {
850        let variants = vec![
851            ExperimentVariant {
852                template_name: "a".into(),
853                weight: 1.0,
854                label: "a".into(),
855            },
856            ExperimentVariant {
857                template_name: "b".into(),
858                weight: 1.0,
859                label: "b".into(),
860            },
861        ];
862        let exp = AbExperiment::new("exp", variants);
863        exp.record_request(0);
864        exp.record_request(0);
865        exp.record_success(0, 100, 0.9);
866        exp.record_success(0, 120, 0.85);
867
868        let metrics = exp.metrics();
869        assert_eq!(metrics[0].requests, 2);
870        assert_eq!(metrics[0].successes, 2);
871        assert!((metrics[0].avg_latency_ms() - 110.0).abs() < 0.01);
872    }
873
874    // -----------------------------------------------------------------------
875    // Block rendering tests
876    // -----------------------------------------------------------------------
877
878    #[test]
879    fn test_if_block_truthy() {
880        let t = PromptTemplate::builder("t_if")
881            .body("Prefix. {{#if show_extra}}Extra content.{{/if}} Suffix.")
882            .build();
883        let mut vars = HashMap::new();
884        vars.insert("show_extra", "true");
885        assert_eq!(t.render(&vars), "Prefix. Extra content. Suffix.");
886    }
887
888    #[test]
889    fn test_if_block_falsy() {
890        let t = PromptTemplate::builder("t_if_false")
891            .body("Prefix. {{#if show_extra}}Extra content.{{/if}} Suffix.")
892            .build();
893        let mut vars = HashMap::new();
894        vars.insert("show_extra", "false");
895        assert_eq!(t.render(&vars), "Prefix.  Suffix.");
896    }
897
898    #[test]
899    fn test_if_block_missing_var_is_falsy() {
900        let t = PromptTemplate::builder("t_if_missing")
901            .body("A{{#if missing}}B{{/if}}C")
902            .build();
903        let vars = HashMap::new();
904        assert_eq!(t.render(&vars), "AC");
905    }
906
907    #[test]
908    fn test_if_block_zero_is_falsy() {
909        let t = PromptTemplate::builder("t_if_zero")
910            .body("{{#if count}}has items{{/if}}")
911            .build();
912        let mut vars = HashMap::new();
913        vars.insert("count", "0");
914        assert_eq!(t.render(&vars), "");
915    }
916
917    #[test]
918    fn test_each_block_basic() {
919        let t = PromptTemplate::builder("t_each")
920            .body("Items: {{#each fruits}}- {{this}}\n{{/each}}")
921            .build();
922        let mut vars = HashMap::new();
923        vars.insert("fruits", "apple, banana, cherry");
924        let result = t.render(&vars);
925        assert!(result.contains("- apple"));
926        assert!(result.contains("- banana"));
927        assert!(result.contains("- cherry"));
928    }
929
930    #[test]
931    fn test_each_block_index() {
932        let t = PromptTemplate::builder("t_each_idx")
933            .body("{{#each items}}{{@index}}:{{this}} {{/each}}")
934            .build();
935        let mut vars = HashMap::new();
936        vars.insert("items", "a,b,c");
937        let result = t.render(&vars);
938        assert!(result.contains("0:a"));
939        assert!(result.contains("1:b"));
940        assert!(result.contains("2:c"));
941    }
942
943    #[test]
944    fn test_each_block_empty_list() {
945        let t = PromptTemplate::builder("t_each_empty")
946            .body("before{{#each items}}{{this}}{{/each}}after")
947            .build();
948        let mut vars = HashMap::new();
949        vars.insert("items", "");
950        assert_eq!(t.render(&vars), "beforeafter");
951    }
952
953    #[test]
954    fn test_if_and_each_combined() {
955        let t = PromptTemplate::builder("t_combined")
956            .body("{{#if show}}List: {{#each items}}{{this}} {{/each}}{{/if}}")
957            .build();
958        let mut vars = HashMap::new();
959        vars.insert("show", "yes");
960        vars.insert("items", "x,y,z");
961        let result = t.render(&vars);
962        assert!(result.contains("List:"));
963        assert!(result.contains("x"));
964        assert!(result.contains("y"));
965        assert!(result.contains("z"));
966    }
967}