use std::collections::HashMap;
use std::sync::Arc;
use anyhow::anyhow;
use anyhow::Result;
pub trait Prompt: Send + Sync {
fn template(&self) -> String;
fn variables(&self) -> Vec<String>;
fn format(&self, input_variables: HashMap<&str, &str>) -> Result<String>;
}
pub struct PromptTemplate {
template: String,
variables: Vec<String>,
}
impl PromptTemplate {
pub fn create(template: &str, variables: Vec<String>) -> Arc<PromptTemplate> {
Arc::new(PromptTemplate {
template: template.to_string(),
variables,
})
}
}
impl Prompt for PromptTemplate {
fn template(&self) -> String {
self.template.clone()
}
fn variables(&self) -> Vec<String> {
self.variables.clone()
}
fn format(&self, input_variables: HashMap<&str, &str>) -> Result<String> {
let mut prompt = self.template();
for (key, value) in input_variables {
if !self.variables().contains(&key.to_string()) {
return Err(anyhow!(
"input variable: '{}' is not in the variables: {:?}",
key,
self.variables()
));
}
let key = format!("{{{}}}", key);
prompt = prompt.replace(&key, value);
}
Ok(prompt)
}
}