1use std::{
40 collections::HashMap,
41 sync::{Arc, RwLock},
42};
43
44#[derive(Debug, Clone)]
54pub struct PromptTemplate {
55 pub name: String,
57 pub version: String,
59 pub system: Option<String>,
61 pub body: String,
63 pub variables: HashMap<String, Option<String>>,
65 pub tags: Vec<String>,
67 pub description: Option<String>,
69}
70
71impl PromptTemplate {
72 pub fn builder(name: impl Into<String>) -> TemplateBuilder {
74 TemplateBuilder::new(name)
75 }
76
77 pub fn render(&self, vars: &HashMap<&str, &str>) -> String {
94 let mut output = render_each_blocks(&self.body, vars);
96
97 output = render_if_blocks(&output, vars);
99
100 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 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 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 pub fn declared_variables(&self) -> Vec<String> {
137 extract_placeholders(&self.body)
138 }
139}
140
141#[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 pub fn var(mut self, name: impl Into<String>) -> Self {
183 self.defaults.insert(name.into(), None);
184 self
185 }
186
187 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 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#[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#[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 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 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 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 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 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 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 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 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#[derive(Debug, Clone)]
397pub struct ExperimentVariant {
398 pub template_name: String,
400 pub weight: f64,
402 pub label: String,
404}
405
406#[derive(Debug, Clone, Default)]
408pub struct VariantMetrics {
409 pub requests: u64,
410 pub successes: u64,
411 pub total_latency_ms: u64,
412 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
439pub 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 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 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 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 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 pub fn metrics(&self) -> Vec<VariantMetrics> {
508 self.metrics.read().unwrap_or_else(|e| e.into_inner()).clone()
509 }
510
511 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 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 let p = 2.0 * standard_normal_cdf(-z.abs());
545 Some(p)
546 }
547
548 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#[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
590fn is_truthy(value: &str) -> bool {
600 !value.is_empty() && value != "false" && value != "0"
601}
602
603fn 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 let tag_end = match output[open_start..].find("}}") {
614 Some(j) => open_start + j + 2,
615 None => break,
616 };
617
618 let condition = output[open_start + 6..tag_end - 2].trim().to_string();
620
621 let close_tag = "{{/if}}";
623 let close_start = match output[tag_end..].find(close_tag) {
624 Some(k) => tag_end + k,
625 None => break, };
627 let close_end = close_start + close_tag.len();
628
629 let body = output[tag_end..close_start].to_string();
631
632 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
646fn 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 let tag_end = match output[open_start..].find("}}") {
661 Some(j) => open_start + j + 2,
662 None => break,
663 };
664
665 let list_key = output[open_start + 8..tag_end - 2].trim().to_string();
667
668 let close_tag = "{{/each}}";
670 let close_start = match output[tag_end..].find(close_tag) {
671 Some(k) => tag_end + k,
672 None => break, };
674 let close_end = close_start + close_tag.len();
675
676 let body_template = output[tag_end..close_start].to_string();
678
679 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
699fn 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
732fn 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#[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 assert_eq!(exp.pick_variant(0.0), 0);
844 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 #[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}