Skip to main content

tokio_prompt_orchestrator/
pipeline_builder.rs

1//! Fluent builder for prompt processing pipelines.
2
3use std::collections::HashMap;
4use std::fmt;
5
6/// Classification of a pipeline stage's role.
7#[derive(Debug, Clone, PartialEq, Eq)]
8pub enum StageKind {
9    /// Normalize input (trim whitespace, change case, etc.).
10    Preprocess,
11    /// Validate input against constraints.
12    Validate,
13    /// Augment input with additional context.
14    Enrich,
15    /// Apply structural transformations (truncate, reformat, etc.).
16    Transform,
17    /// Apply finishing touches to output.
18    Postprocess,
19}
20
21/// A single processing stage in a pipeline.
22#[derive(Debug, Clone)]
23pub struct PipelineStage {
24    /// Unique human-readable name for this stage.
25    pub name: String,
26    /// What kind of processing this stage performs.
27    pub kind: StageKind,
28    /// Whether this stage participates in pipeline runs.
29    pub enabled: bool,
30    /// Arbitrary string→string configuration for the stage.
31    pub config: HashMap<String, String>,
32}
33
34/// Result produced by running a single stage.
35#[derive(Debug, Clone)]
36pub struct StageResult {
37    /// Name of the stage that produced this result.
38    pub stage_name: String,
39    /// Whether the stage altered the input string.
40    pub modified: bool,
41    /// The output string after stage processing.
42    pub output: String,
43    /// Wall-clock time taken by this stage, in microseconds.
44    pub elapsed_us: u64,
45    /// Arbitrary metadata emitted by the stage.
46    pub metadata: HashMap<String, String>,
47}
48
49/// Errors that can occur during pipeline construction or execution.
50#[derive(Debug)]
51pub enum PipelineError {
52    /// A specific stage produced an error.
53    StageError {
54        /// Name of the failing stage.
55        stage: String,
56        /// Human-readable reason for the failure.
57        reason: String,
58    },
59    /// Pipeline or stage configuration is invalid.
60    InvalidConfig(String),
61    /// [`PipelineBuilder::build`] was called with no stages added.
62    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/// Full configuration for a [`Pipeline`].
78#[derive(Debug, Clone)]
79pub struct PipelineConfig {
80    /// Ordered list of stages.
81    pub stages: Vec<PipelineStage>,
82    /// When `true`, the pipeline aborts on the first stage error.
83    pub fail_fast: bool,
84    /// Budget for the entire pipeline run in microseconds (not enforced in this
85    /// implementation — stored for future use).
86    pub max_total_us: u64,
87}
88
89/// Fluent builder for [`Pipeline`].
90pub struct PipelineBuilder {
91    stages: Vec<PipelineStage>,
92    fail_fast: bool,
93    max_total_us: u64,
94}
95
96impl PipelineBuilder {
97    /// Create a new empty builder.
98    pub fn new() -> Self {
99        Self {
100            stages: Vec::new(),
101            fail_fast: false,
102            max_total_us: u64::MAX,
103        }
104    }
105
106    /// Append a new stage with the given name and kind.
107    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    /// Set a configuration key/value pair on the most-recently added stage
118    /// named `name`.  Does nothing if no stage with that name exists yet.
119    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    /// Disable the stage named `name`.  Does nothing if no such stage exists.
127    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    /// Set whether the pipeline should abort on the first stage error.
135    pub fn fail_fast(mut self, v: bool) -> Self {
136        self.fail_fast = v;
137        self
138    }
139
140    /// Consume the builder and produce a [`Pipeline`].
141    /// Returns [`PipelineError::EmptyPipeline`] if no stages were added.
142    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
162/// A constructed, immutable processing pipeline.
163pub struct Pipeline {
164    /// The resolved configuration.
165    pub config: PipelineConfig,
166}
167
168impl Pipeline {
169    /// Run the pipeline on `input`, returning ordered results for each enabled stage.
170    ///
171    /// Stage semantics:
172    /// - **Preprocess**: trim whitespace; if `config["lowercase"] == "true"`, convert to lowercase.
173    /// - **Validate**: return a [`PipelineError::StageError`] if the current text is empty.
174    /// - **Enrich**: prepend `config["prefix"]` (if set).
175    /// - **Transform**: truncate to `config["max_len"]` characters (if set and parseable).
176    /// - **Postprocess**: append `config["suffix"]` (if set).
177    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, &current);
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                    // On non-fail-fast, record the pre-stage text as output and continue.
206                    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
225/// Apply a single stage's transformation to `text`.
226fn 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
275/// Return the output string of the last [`StageResult`] in `results`, or `None`
276/// if the slice is empty.
277pub 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        // Empty input should fail Validate
314        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        // Only "pre" ran
335        assert_eq!(results.len(), 1);
336        assert_eq!(results[0].output, "hello"); // lowercase applied, NOT truncated
337    }
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}