tokio_prompt_orchestrator/
pipeline_builder.rs1use std::collections::HashMap;
4use std::fmt;
5
6#[derive(Debug, Clone, PartialEq, Eq)]
8pub enum StageKind {
9 Preprocess,
11 Validate,
13 Enrich,
15 Transform,
17 Postprocess,
19}
20
21#[derive(Debug, Clone)]
23pub struct PipelineStage {
24 pub name: String,
26 pub kind: StageKind,
28 pub enabled: bool,
30 pub config: HashMap<String, String>,
32}
33
34#[derive(Debug, Clone)]
36pub struct StageResult {
37 pub stage_name: String,
39 pub modified: bool,
41 pub output: String,
43 pub elapsed_us: u64,
45 pub metadata: HashMap<String, String>,
47}
48
49#[derive(Debug)]
51pub enum PipelineError {
52 StageError {
54 stage: String,
56 reason: String,
58 },
59 InvalidConfig(String),
61 EmptyPipeline,
63}
64
65impl fmt::Display for PipelineError {
66 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
67 match self {
68 PipelineError::StageError { stage, reason } => {
69 write!(f, "stage '{}' failed: {}", stage, reason)
70 }
71 PipelineError::InvalidConfig(msg) => write!(f, "invalid config: {}", msg),
72 PipelineError::EmptyPipeline => write!(f, "pipeline has no stages"),
73 }
74 }
75}
76
77#[derive(Debug, Clone)]
79pub struct PipelineConfig {
80 pub stages: Vec<PipelineStage>,
82 pub fail_fast: bool,
84 pub max_total_us: u64,
87}
88
89pub struct PipelineBuilder {
91 stages: Vec<PipelineStage>,
92 fail_fast: bool,
93 max_total_us: u64,
94}
95
96impl PipelineBuilder {
97 pub fn new() -> Self {
99 Self {
100 stages: Vec::new(),
101 fail_fast: false,
102 max_total_us: u64::MAX,
103 }
104 }
105
106 pub fn add_stage(mut self, name: &str, kind: StageKind) -> Self {
108 self.stages.push(PipelineStage {
109 name: name.to_string(),
110 kind,
111 enabled: true,
112 config: HashMap::new(),
113 });
114 self
115 }
116
117 pub fn configure_stage(mut self, name: &str, key: &str, value: &str) -> Self {
120 if let Some(stage) = self.stages.iter_mut().rev().find(|s| s.name == name) {
121 stage.config.insert(key.to_string(), value.to_string());
122 }
123 self
124 }
125
126 pub fn disable_stage(mut self, name: &str) -> Self {
128 if let Some(stage) = self.stages.iter_mut().find(|s| s.name == name) {
129 stage.enabled = false;
130 }
131 self
132 }
133
134 pub fn fail_fast(mut self, v: bool) -> Self {
136 self.fail_fast = v;
137 self
138 }
139
140 pub fn build(self) -> Result<Pipeline, PipelineError> {
143 if self.stages.is_empty() {
144 return Err(PipelineError::EmptyPipeline);
145 }
146 Ok(Pipeline {
147 config: PipelineConfig {
148 stages: self.stages,
149 fail_fast: self.fail_fast,
150 max_total_us: self.max_total_us,
151 },
152 })
153 }
154}
155
156impl Default for PipelineBuilder {
157 fn default() -> Self {
158 Self::new()
159 }
160}
161
162pub struct Pipeline {
164 pub config: PipelineConfig,
166}
167
168impl Pipeline {
169 pub fn run(&self, input: &str) -> Result<Vec<StageResult>, PipelineError> {
178 let mut current = input.to_string();
179 let mut results = Vec::new();
180
181 for stage in &self.config.stages {
182 if !stage.enabled {
183 continue;
184 }
185
186 let before = current.clone();
187 let stage_output = apply_stage(stage, ¤t);
188
189 match stage_output {
190 Ok(text) => {
191 let modified = text != before;
192 current = text.clone();
193 results.push(StageResult {
194 stage_name: stage.name.clone(),
195 modified,
196 output: text,
197 elapsed_us: 0,
198 metadata: HashMap::new(),
199 });
200 }
201 Err(e) => {
202 if self.config.fail_fast {
203 return Err(e);
204 }
205 results.push(StageResult {
207 stage_name: stage.name.clone(),
208 modified: false,
209 output: current.clone(),
210 elapsed_us: 0,
211 metadata: {
212 let mut m = HashMap::new();
213 m.insert("error".to_string(), e.to_string());
214 m
215 },
216 });
217 }
218 }
219 }
220
221 Ok(results)
222 }
223}
224
225fn apply_stage(stage: &PipelineStage, text: &str) -> Result<String, PipelineError> {
227 match stage.kind {
228 StageKind::Preprocess => {
229 let mut out = text.trim().to_string();
230 if stage.config.get("lowercase").map(|v| v == "true").unwrap_or(false) {
231 out = out.to_lowercase();
232 }
233 Ok(out)
234 }
235 StageKind::Validate => {
236 if text.trim().is_empty() {
237 Err(PipelineError::StageError {
238 stage: stage.name.clone(),
239 reason: "input is empty".to_string(),
240 })
241 } else {
242 Ok(text.to_string())
243 }
244 }
245 StageKind::Enrich => {
246 if let Some(prefix) = stage.config.get("prefix") {
247 Ok(format!("{}{}", prefix, text))
248 } else {
249 Ok(text.to_string())
250 }
251 }
252 StageKind::Transform => {
253 if let Some(max_len_str) = stage.config.get("max_len") {
254 match max_len_str.parse::<usize>() {
255 Ok(max_len) => Ok(text.chars().take(max_len).collect()),
256 Err(_) => Err(PipelineError::InvalidConfig(format!(
257 "stage '{}': max_len '{}' is not a valid usize",
258 stage.name, max_len_str
259 ))),
260 }
261 } else {
262 Ok(text.to_string())
263 }
264 }
265 StageKind::Postprocess => {
266 if let Some(suffix) = stage.config.get("suffix") {
267 Ok(format!("{}{}", text, suffix))
268 } else {
269 Ok(text.to_string())
270 }
271 }
272 }
273}
274
275pub fn last_output(results: &[StageResult]) -> Option<&str> {
278 results.last().map(|r| r.output.as_str())
279}
280
281#[cfg(test)]
282mod tests {
283 use super::*;
284
285 #[test]
286 fn test_three_stages_all_run() {
287 let pipeline = PipelineBuilder::new()
288 .add_stage("pre", StageKind::Preprocess)
289 .configure_stage("pre", "lowercase", "true")
290 .add_stage("validate", StageKind::Validate)
291 .add_stage("post", StageKind::Postprocess)
292 .configure_stage("post", "suffix", "!")
293 .build()
294 .unwrap();
295
296 let results = pipeline.run(" HELLO WORLD ").unwrap();
297 assert_eq!(results.len(), 3);
298 assert_eq!(results[0].output, "hello world");
299 assert_eq!(results[1].output, "hello world");
300 assert_eq!(results[2].output, "hello world!");
301 }
302
303 #[test]
304 fn test_fail_fast_stops_on_error() {
305 let pipeline = PipelineBuilder::new()
306 .add_stage("validate", StageKind::Validate)
307 .add_stage("post", StageKind::Postprocess)
308 .configure_stage("post", "suffix", "!")
309 .fail_fast(true)
310 .build()
311 .unwrap();
312
313 let err = pipeline.run(" ");
315 assert!(err.is_err());
316 match err.unwrap_err() {
317 PipelineError::StageError { stage, .. } => assert_eq!(stage, "validate"),
318 other => panic!("unexpected error: {}", other),
319 }
320 }
321
322 #[test]
323 fn test_disabled_stage_skipped() {
324 let pipeline = PipelineBuilder::new()
325 .add_stage("pre", StageKind::Preprocess)
326 .configure_stage("pre", "lowercase", "true")
327 .add_stage("skip_me", StageKind::Transform)
328 .configure_stage("skip_me", "max_len", "3")
329 .disable_stage("skip_me")
330 .build()
331 .unwrap();
332
333 let results = pipeline.run("HELLO").unwrap();
334 assert_eq!(results.len(), 1);
336 assert_eq!(results[0].output, "hello"); }
338
339 #[test]
340 fn test_empty_pipeline_error() {
341 let err = PipelineBuilder::new().build();
342 assert!(matches!(err, Err(PipelineError::EmptyPipeline)));
343 }
344
345 #[test]
346 fn test_lowercase_transform() {
347 let pipeline = PipelineBuilder::new()
348 .add_stage("pre", StageKind::Preprocess)
349 .configure_stage("pre", "lowercase", "true")
350 .build()
351 .unwrap();
352
353 let results = pipeline.run("FoO BaR").unwrap();
354 assert_eq!(last_output(&results), Some("foo bar"));
355 }
356
357 #[test]
358 fn test_enrich_prepend() {
359 let pipeline = PipelineBuilder::new()
360 .add_stage("enrich", StageKind::Enrich)
361 .configure_stage("enrich", "prefix", ">>")
362 .build()
363 .unwrap();
364
365 let results = pipeline.run("text").unwrap();
366 assert_eq!(last_output(&results), Some(">>text"));
367 }
368
369 #[test]
370 fn test_transform_truncate() {
371 let pipeline = PipelineBuilder::new()
372 .add_stage("trunc", StageKind::Transform)
373 .configure_stage("trunc", "max_len", "4")
374 .build()
375 .unwrap();
376
377 let results = pipeline.run("abcdefgh").unwrap();
378 assert_eq!(last_output(&results), Some("abcd"));
379 }
380}