1use std::collections::BTreeMap;
19use std::fs;
20use std::path::{Path, PathBuf};
21
22use serde::Deserialize;
23use thiserror::Error;
24use workload_spec::ImageRef;
25
26use velveteen::TaskRuntime;
27
28pub const ENV_TRANSFORM_IN_0: &str = "YAH_TRANSFORM_IN_0";
30
31pub const ENV_TRANSFORM_OUT: &str = "YAH_TRANSFORM_OUT";
33
34#[derive(Debug, Clone, PartialEq, Eq, Deserialize)]
45pub struct TransformRecipe {
46 pub name: String,
47 pub label: String,
48 pub placement: RecipePlacement,
49 pub image: ImageRef,
50 #[serde(default)]
51 pub steps: Vec<RecipeStep>,
52}
53
54#[derive(Debug, Clone, PartialEq, Eq, Deserialize)]
67pub struct RecipePlacement {
68 pub location: RecipeLocation,
69 pub runtime: TaskRuntime,
70 #[serde(default)]
71 pub platform: Option<String>,
72}
73
74#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize)]
76#[serde(rename_all = "snake_case")]
77pub enum RecipeLocation {
78 Local,
79}
80
81#[derive(Debug, Clone, PartialEq, Eq, Deserialize)]
85pub struct RecipeStep {
86 pub name: String,
87 pub argv: Vec<String>,
88 #[serde(default)]
90 pub timeout: u64,
91}
92
93#[derive(Error, Debug)]
94pub enum RecipeError {
95 #[error("recipe {name:?} not found at {path}")]
96 NotFound { name: String, path: PathBuf },
97 #[error("reading {path}: {source}")]
98 Io {
99 path: PathBuf,
100 #[source]
101 source: std::io::Error,
102 },
103 #[error("parsing recipe {path}: {source}")]
104 Parse {
105 path: PathBuf,
106 #[source]
107 source: toml::de::Error,
108 },
109 #[error(
110 "recipe {name:?} at {path} uses a bare-tag image; recipe images must \
111 be digest-pinned (e.g. `image = \"...:v1@sha256:<hex>\"`) for \
112 reproducibility (W164)"
113 )]
114 ImageNotPinned { name: String, path: PathBuf },
115}
116
117pub struct TransformRecipeLoader {
124 transforms_dir: PathBuf,
125}
126
127impl TransformRecipeLoader {
128 pub fn new(transforms_dir: impl AsRef<Path>) -> Self {
131 Self {
132 transforms_dir: transforms_dir.as_ref().to_path_buf(),
133 }
134 }
135
136 pub fn recipe_path(&self, name: &str) -> PathBuf {
138 self.transforms_dir.join(format!("{name}.toml"))
139 }
140
141 pub fn list_all(&self) -> Result<Vec<String>, RecipeError> {
147 let entries = match fs::read_dir(&self.transforms_dir) {
148 Ok(e) => e,
149 Err(source) if source.kind() == std::io::ErrorKind::NotFound => return Ok(Vec::new()),
150 Err(source) => {
151 return Err(RecipeError::Io {
152 path: self.transforms_dir.clone(),
153 source,
154 })
155 }
156 };
157 let mut names: Vec<String> = entries
158 .filter_map(|e| e.ok())
159 .map(|e| e.path())
160 .filter(|p| p.extension().is_some_and(|ext| ext == "toml"))
161 .filter_map(|p| p.file_stem().map(|s| s.to_string_lossy().into_owned()))
162 .collect();
163 names.sort();
164 Ok(names)
165 }
166
167 pub fn load(&self, name: &str) -> Result<TransformRecipe, RecipeError> {
174 let path = self.recipe_path(name);
175 if !path.exists() {
176 return Err(RecipeError::NotFound { name: name.to_string(), path });
177 }
178 self.load_from_path(&path)
179 }
180
181 pub fn load_from_path(&self, path: &Path) -> Result<TransformRecipe, RecipeError> {
185 let content = fs::read_to_string(path).map_err(|source| RecipeError::Io {
186 path: path.to_path_buf(),
187 source,
188 })?;
189 let recipe: TransformRecipe = toml::from_str(&content).map_err(|source| RecipeError::Parse {
190 path: path.to_path_buf(),
191 source,
192 })?;
193 if recipe.image.digest.is_empty() {
197 return Err(RecipeError::ImageNotPinned {
198 name: recipe.name.clone(),
199 path: path.to_path_buf(),
200 });
201 }
202 Ok(recipe)
203 }
204}
205
206pub fn substitute_argv(template: &[String], params: &BTreeMap<String, String>) -> Vec<String> {
218 template.iter().map(|elem| substitute_one(elem, params)).collect()
219}
220
221fn substitute_one(elem: &str, params: &BTreeMap<String, String>) -> String {
222 let mut out = String::with_capacity(elem.len());
223 let mut rest = elem;
224 while let Some(start) = rest.find("{{") {
225 out.push_str(&rest[..start]);
226 let after = &rest[start + 2..];
227 let Some(end) = after.find("}}") else {
228 out.push_str("{{");
230 out.push_str(after);
231 return out;
232 };
233 let key = after[..end].trim();
234 if let Some(val) = params.get(key) {
235 out.push_str(val);
236 } else {
237 out.push_str("{{");
239 out.push_str(&after[..end]);
240 out.push_str("}}");
241 }
242 rest = &after[end + 2..];
243 }
244 out.push_str(rest);
245 out
246}
247
248#[cfg(test)]
249mod tests {
250 use super::*;
251 use tempfile::tempdir;
252
253 const HASH_64: &str = "abcdef1234567890abcdef1234567890abcdef1234567890abcdef1234567890";
254
255 fn whisper_quantize_toml() -> String {
256 format!(
260 r#"
261name = "whisper-quantize"
262label = "Quantize a whisper GGML model"
263image = "ghcr.io/ggerganov/whisper.cpp:v1.7.4@sha256:{HASH_64}"
264
265[placement]
266location = "local"
267runtime = "container"
268
269[[steps]]
270name = "quantize"
271argv = ["./quantize", "{{{{YAH_TRANSFORM_IN_0}}}}", "{{{{YAH_TRANSFORM_OUT}}}}", "{{{{quant}}}}"]
272timeout = 600
273"#
274 )
275 }
276
277 #[test]
278 fn loads_sample_recipe_round_trip() {
279 let dir = tempdir().unwrap();
280 let transforms = dir.path().join("transforms");
281 fs::create_dir_all(&transforms).unwrap();
282 fs::write(
283 transforms.join("whisper-quantize.toml"),
284 whisper_quantize_toml(),
285 )
286 .unwrap();
287
288 let loader = TransformRecipeLoader::new(&transforms);
289 let recipe = loader.load("whisper-quantize").expect("load recipe");
290 assert_eq!(recipe.name, "whisper-quantize");
291 assert_eq!(recipe.placement.location, RecipeLocation::Local);
292 assert_eq!(recipe.placement.runtime, TaskRuntime::Container);
293 assert_eq!(recipe.image.registry, "ghcr.io");
294 assert_eq!(recipe.image.repository, "ggerganov/whisper.cpp");
295 assert_eq!(recipe.image.tag, "v1.7.4");
296 assert_eq!(recipe.image.digest, format!("sha256:{HASH_64}"));
297 assert_eq!(recipe.steps.len(), 1);
298 assert_eq!(recipe.steps[0].name, "quantize");
299 assert_eq!(recipe.steps[0].timeout, 600);
300 assert_eq!(
301 recipe.steps[0].argv,
302 vec![
303 "./quantize",
304 "{{YAH_TRANSFORM_IN_0}}",
305 "{{YAH_TRANSFORM_OUT}}",
306 "{{quant}}",
307 ]
308 );
309 }
310
311 #[test]
312 fn list_all_returns_sorted_stems_and_ignores_non_toml() {
313 let dir = tempdir().unwrap();
314 let transforms = dir.path().join("transforms");
315 fs::create_dir_all(&transforms).unwrap();
316 fs::write(transforms.join("zeta.toml"), "").unwrap();
317 fs::write(transforms.join("alpha.toml"), "").unwrap();
318 fs::write(transforms.join("README.md"), "ignore me").unwrap();
319
320 let loader = TransformRecipeLoader::new(&transforms);
321 let names = loader.list_all().expect("list_all ok");
322 assert_eq!(names, vec!["alpha".to_string(), "zeta".to_string()]);
323 }
324
325 #[test]
326 fn list_all_missing_dir_is_empty_not_error() {
327 let dir = tempdir().unwrap();
328 let loader = TransformRecipeLoader::new(dir.path().join("does-not-exist"));
329 assert_eq!(loader.list_all().expect("ok"), Vec::<String>::new());
330 }
331
332 #[test]
333 fn rejects_recipe_with_bare_tag_string_image() {
334 let dir = tempdir().unwrap();
335 let transforms = dir.path().join("transforms");
336 fs::create_dir_all(&transforms).unwrap();
337 let bad = r#"
338name = "bare-tag"
339label = "no pin"
340image = "node:20"
341
342[placement]
343location = "local"
344runtime = "container"
345
346[[steps]]
347name = "noop"
348argv = ["true"]
349"#;
350 fs::write(transforms.join("bare-tag.toml"), bad).unwrap();
351 let loader = TransformRecipeLoader::new(&transforms);
352 let err = loader.load("bare-tag").expect_err("bare-tag must reject");
353 assert!(matches!(err, RecipeError::Parse { .. }), "got {err:?}");
355 }
356
357 #[test]
358 fn rejects_recipe_with_struct_image_missing_digest() {
359 let dir = tempdir().unwrap();
360 let transforms = dir.path().join("transforms");
361 fs::create_dir_all(&transforms).unwrap();
362 let bad = r#"
363name = "struct-bare"
364label = "struct-form bare tag"
365
366[placement]
367location = "local"
368runtime = "container"
369
370[image]
371registry = "ghcr.io"
372repository = "foo/bar"
373tag = "v1"
374
375[[steps]]
376name = "noop"
377argv = ["true"]
378"#;
379 fs::write(transforms.join("struct-bare.toml"), bad).unwrap();
380 let loader = TransformRecipeLoader::new(&transforms);
381 let err = loader
382 .load("struct-bare")
383 .expect_err("struct-form bare tag must reject");
384 assert!(matches!(err, RecipeError::Parse { .. }), "got {err:?}");
389 }
390
391 #[test]
392 fn missing_recipe_reports_path() {
393 let dir = tempdir().unwrap();
394 let loader = TransformRecipeLoader::new(dir.path().join("transforms"));
395 let err = loader.load("nope").expect_err("missing recipe");
396 match err {
397 RecipeError::NotFound { name, .. } => assert_eq!(name, "nope"),
398 other => panic!("expected NotFound, got {other:?}"),
399 }
400 }
401
402 #[test]
403 fn substitute_argv_replaces_known_placeholders() {
404 let template = vec![
405 "./tool".to_string(),
406 "{{YAH_TRANSFORM_IN_0}}".to_string(),
407 "{{YAH_TRANSFORM_OUT}}".to_string(),
408 "--mode={{quant}}".to_string(),
409 ];
410 let mut params = BTreeMap::new();
411 params.insert(ENV_TRANSFORM_IN_0.into(), "/cache/fetch/abc.bin".into());
412 params.insert(ENV_TRANSFORM_OUT.into(), "/tmp/out.bin".into());
413 params.insert("quant".into(), "q5_1".into());
414 let resolved = substitute_argv(&template, ¶ms);
415 assert_eq!(
416 resolved,
417 vec![
418 "./tool",
419 "/cache/fetch/abc.bin",
420 "/tmp/out.bin",
421 "--mode=q5_1",
422 ]
423 );
424 }
425
426 #[test]
427 fn substitute_argv_preserves_unknown_placeholders() {
428 let template = vec!["{{unknown}}".to_string()];
429 let resolved = substitute_argv(&template, &BTreeMap::new());
430 assert_eq!(resolved, vec!["{{unknown}}".to_string()]);
431 }
432
433 #[test]
434 fn substitute_argv_does_not_split_values_with_spaces() {
435 let template = vec!["{{flag}}".to_string()];
438 let mut params = BTreeMap::new();
439 params.insert("flag".into(), "--a --b --c".into());
440 let resolved = substitute_argv(&template, ¶ms);
441 assert_eq!(resolved, vec!["--a --b --c"]);
442 assert_eq!(resolved.len(), 1, "must not split on spaces");
443 }
444
445 #[test]
446 fn substitute_argv_handles_unterminated_braces() {
447 let template = vec!["{{never_closed".to_string()];
448 let resolved = substitute_argv(&template, &BTreeMap::new());
449 assert_eq!(resolved, vec!["{{never_closed".to_string()]);
450 }
451
452 #[test]
453 fn substitute_argv_trims_whitespace_in_key() {
454 let template = vec!["{{ key }}".to_string()];
455 let mut params = BTreeMap::new();
456 params.insert("key".into(), "value".into());
457 let resolved = substitute_argv(&template, ¶ms);
458 assert_eq!(resolved, vec!["value".to_string()]);
459 }
460}