rskit_ai/prompt/
template.rs1use std::collections::{BTreeMap, BTreeSet};
4
5use semver::Version;
6use serde::{Deserialize, Serialize};
7use thiserror::Error;
8
9use super::render::{placeholders, render};
10
11#[derive(Debug, Error)]
13#[non_exhaustive]
14pub enum PromptError {
15 #[error("missing prompt variable {0:?}")]
17 MissingVariable(String),
18 #[error("missing required prompt field {0}")]
20 MissingField(&'static str),
21 #[error("invalid prompt version {version:?}: {source}")]
23 InvalidVersion {
24 version: String,
26 #[source]
28 source: semver::Error,
29 },
30 #[error("prompt already registered: {name}@{version}")]
32 AlreadyRegistered {
33 name: String,
35 version: Version,
37 },
38 #[error("prompt not found: {name}@{version}")]
40 NotFound {
41 name: String,
43 version: Version,
45 },
46 #[error("prompt not found: {0}")]
48 NameNotFound(String),
49}
50
51pub type RenderContext = BTreeMap<String, serde_json::Value>;
53
54#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
56#[serde(rename_all = "snake_case")]
57#[non_exhaustive]
58pub enum VariableType {
59 String,
61 Number,
63 Boolean,
65 Object,
67 Array,
69 #[default]
71 Any,
72}
73
74#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
76pub struct VariableDecl {
77 pub name: String,
79 #[serde(default, rename = "type")]
81 pub kind: VariableType,
82 #[serde(default = "default_required")]
84 pub required: bool,
85 #[serde(default, skip_serializing_if = "Option::is_none")]
87 pub default: Option<serde_json::Value>,
88}
89
90const fn default_required() -> bool {
91 true
92}
93
94#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
96pub struct PromptTemplate {
97 pub name: String,
99 pub version: Version,
101 pub template: String,
103 #[serde(default, skip_serializing_if = "Vec::is_empty")]
105 pub variables: Vec<VariableDecl>,
106 #[serde(default, skip_serializing_if = "Option::is_none")]
108 pub output_schema: Option<serde_json::Value>,
109 #[serde(default, skip_serializing_if = "String::is_empty")]
111 pub description: String,
112}
113
114impl PromptTemplate {
115 pub fn render(&self, context: &RenderContext) -> Result<String, PromptError> {
117 let mut values = context.clone();
118 for variable in &self.variables {
119 if !values.contains_key(&variable.name) {
120 if let Some(default) = &variable.default {
121 values.insert(variable.name.clone(), default.clone());
122 } else if variable.required {
123 return Err(PromptError::MissingVariable(variable.name.clone()));
124 }
125 }
126 }
127 render(&self.template, &values)
128 }
129}
130
131pub type Template = PromptTemplate;
133
134#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
136#[serde(rename_all = "snake_case")]
137pub enum ValidationFindingKind {
138 MissingVariable,
140 UnusedVariable,
142}
143
144#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
146pub struct ValidationFinding {
147 pub kind: ValidationFindingKind,
149 pub variable: String,
151}
152
153#[must_use]
155pub fn validate(template: &PromptTemplate) -> Vec<ValidationFinding> {
156 let used = placeholders(&template.template);
157 let declared = template
158 .variables
159 .iter()
160 .map(|variable| variable.name.clone())
161 .collect::<BTreeSet<_>>();
162 let mut findings = Vec::new();
163 for variable in used.difference(&declared) {
164 findings.push(ValidationFinding {
165 kind: ValidationFindingKind::MissingVariable,
166 variable: variable.clone(),
167 });
168 }
169 for variable in declared.difference(&used) {
170 findings.push(ValidationFinding {
171 kind: ValidationFindingKind::UnusedVariable,
172 variable: variable.clone(),
173 });
174 }
175 findings
176}