use regex::Regex;
use std::collections::HashMap;
use std::sync::LazyLock;
static VARIABLE_RE: LazyLock<Regex> = LazyLock::new(|| Regex::new(r"\{(\w+)\}").unwrap());
pub struct PromptTemplate {
template: String,
}
impl PromptTemplate {
pub fn new(template: impl Into<String>) -> Self {
Self {
template: template.into(),
}
}
pub fn format(&self, variables: &HashMap<&str, &str>) -> Result<String, String> {
let mut result = String::with_capacity(self.template.len());
let mut last_end = 0;
for cap in VARIABLE_RE.captures_iter(&self.template) {
let var_match = cap.get(0).unwrap();
let var_name = cap.get(1).unwrap().as_str();
result.push_str(&self.template[last_end..var_match.start()]);
if let Some(value) = variables.get(var_name) {
result.push_str(value);
} else {
return Err(format!("Missing variable: {}", var_name));
}
last_end = var_match.end();
}
result.push_str(&self.template[last_end..]);
Ok(result)
}
pub fn variables(&self) -> Vec<String> {
VARIABLE_RE
.captures_iter(&self.template)
.map(|cap| cap.get(1).unwrap().as_str().to_string())
.collect()
}
pub fn template(&self) -> &str {
&self.template
}
}
impl std::fmt::Display for PromptTemplate {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "{}", self.template)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_basic_template() {
let template = PromptTemplate::new("你好,{name}!");
let mut vars = HashMap::new();
vars.insert("name", "小明");
let result = template.format(&vars).unwrap();
assert_eq!(result, "你好,小明!");
}
#[test]
fn test_multiple_variables() {
let template = PromptTemplate::new("{greeting},{name}!今天是{day}。");
let mut vars = HashMap::new();
vars.insert("greeting", "早上好");
vars.insert("name", "小红");
vars.insert("day", "星期一");
let result = template.format(&vars).unwrap();
assert_eq!(result, "早上好,小红!今天是星期一。");
}
#[test]
fn test_missing_variable() {
let template = PromptTemplate::new("你好,{name}!今天是{day}。");
let mut vars = HashMap::new();
vars.insert("name", "小明");
let result = template.format(&vars);
assert!(result.is_err());
assert!(result.unwrap_err().contains("day"));
}
#[test]
fn test_get_variables() {
let template = PromptTemplate::new("{a}, {b}, {c}");
let vars = template.variables();
assert_eq!(vars, vec!["a", "b", "c"]);
}
}