use std::collections::{BTreeMap, BTreeSet};
use semver::Version;
use serde::{Deserialize, Serialize};
use thiserror::Error;
use super::render::{placeholders, render};
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum PromptError {
#[error("missing prompt variable {0:?}")]
MissingVariable(String),
#[error("missing required prompt field {0}")]
MissingField(&'static str),
#[error("invalid prompt version {version:?}: {source}")]
InvalidVersion {
version: String,
#[source]
source: semver::Error,
},
#[error("prompt already registered: {name}@{version}")]
AlreadyRegistered {
name: String,
version: Version,
},
#[error("prompt not found: {name}@{version}")]
NotFound {
name: String,
version: Version,
},
#[error("prompt not found: {0}")]
NameNotFound(String),
}
pub type RenderContext = BTreeMap<String, serde_json::Value>;
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
#[non_exhaustive]
pub enum VariableType {
String,
Number,
Boolean,
Object,
Array,
#[default]
Any,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct VariableDecl {
pub name: String,
#[serde(default, rename = "type")]
pub kind: VariableType,
#[serde(default = "default_required")]
pub required: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub default: Option<serde_json::Value>,
}
const fn default_required() -> bool {
true
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct PromptTemplate {
pub name: String,
pub version: Version,
pub template: String,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub variables: Vec<VariableDecl>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_schema: Option<serde_json::Value>,
#[serde(default, skip_serializing_if = "String::is_empty")]
pub description: String,
}
impl PromptTemplate {
pub fn render(&self, context: &RenderContext) -> Result<String, PromptError> {
let mut values = context.clone();
for variable in &self.variables {
if !values.contains_key(&variable.name) {
if let Some(default) = &variable.default {
values.insert(variable.name.clone(), default.clone());
} else if variable.required {
return Err(PromptError::MissingVariable(variable.name.clone()));
}
}
}
render(&self.template, &values)
}
}
pub type Template = PromptTemplate;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum ValidationFindingKind {
MissingVariable,
UnusedVariable,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ValidationFinding {
pub kind: ValidationFindingKind,
pub variable: String,
}
#[must_use]
pub fn validate(template: &PromptTemplate) -> Vec<ValidationFinding> {
let used = placeholders(&template.template);
let declared = template
.variables
.iter()
.map(|variable| variable.name.clone())
.collect::<BTreeSet<_>>();
let mut findings = Vec::new();
for variable in used.difference(&declared) {
findings.push(ValidationFinding {
kind: ValidationFindingKind::MissingVariable,
variable: variable.clone(),
});
}
for variable in declared.difference(&used) {
findings.push(ValidationFinding {
kind: ValidationFindingKind::UnusedVariable,
variable: variable.clone(),
});
}
findings
}