Skip to main content

tokio_prompt_orchestrator/
pipeline.rs

1//! # Prompt Pipeline
2//!
3//! A composable, ordered sequence of text-transformation stages.  Each stage
4//! receives the previous stage's output as its input.  Stages are async and
5//! may be composed with [`PipelineBuilder`].
6//!
7//! ## Example
8//!
9//! ```rust
10//! use tokio_prompt_orchestrator::pipeline::{PipelineBuilder, TrimStage, TruncateStage};
11//!
12//! # #[tokio::main]
13//! # async fn main() {
14//! let pipeline = PipelineBuilder::new()
15//!     .add(TrimStage)
16//!     .add(TruncateStage { max_chars: 100 })
17//!     .build();
18//!
19//! let result = pipeline.run("  hello world  ".to_string()).await.unwrap();
20//! assert_eq!(result.output, "hello world");
21//! # }
22//! ```
23
24use async_trait::async_trait;
25use regex::Regex;
26use std::time::Instant;
27use thiserror::Error;
28
29// ── Error ─────────────────────────────────────────────────────────────────────
30
31/// Errors produced by pipeline stages.
32#[derive(Debug, Error)]
33pub enum PipelineError {
34    /// A stage produced an I/O or processing error.
35    #[error("stage '{stage}' failed: {message}")]
36    StageError {
37        /// Name of the stage that failed.
38        stage: String,
39        /// Human-readable description of the failure.
40        message: String,
41    },
42
43    /// A regex pattern provided to [`RegexReplaceStage`] was invalid.
44    #[error("invalid regex pattern '{pattern}': {source}")]
45    InvalidRegex {
46        /// The offending pattern string.
47        pattern: String,
48        /// The underlying regex compile error.
49        #[source]
50        source: regex::Error,
51    },
52}
53
54// ── Trait ─────────────────────────────────────────────────────────────────────
55
56/// A single processing stage in a [`Pipeline`].
57///
58/// Implementors transform an input string and return the (possibly mutated)
59/// result.  Stages are async to allow future network-based transformations
60/// (e.g. embedding lookups, remote filters) without blocking the Tokio runtime.
61#[async_trait]
62pub trait PipelineStage: Send + Sync {
63    /// Process `input` and return the transformed string.
64    async fn process(&self, input: String) -> Result<String, PipelineError>;
65
66    /// A short human-readable name used in error messages and statistics.
67    fn name(&self) -> &str;
68}
69
70// ── Built-in stages ───────────────────────────────────────────────────────────
71
72/// Strips leading and trailing ASCII whitespace from the input.
73pub struct TrimStage;
74
75#[async_trait]
76impl PipelineStage for TrimStage {
77    async fn process(&self, input: String) -> Result<String, PipelineError> {
78        Ok(input.trim().to_string())
79    }
80
81    fn name(&self) -> &str {
82        "TrimStage"
83    }
84}
85
86/// Truncates the input to at most `max_chars` characters, breaking at the last
87/// word boundary that fits within the limit.
88///
89/// If `max_chars` is zero the output is always an empty string.
90pub struct TruncateStage {
91    /// Maximum number of characters to keep.
92    pub max_chars: usize,
93}
94
95#[async_trait]
96impl PipelineStage for TruncateStage {
97    async fn process(&self, input: String) -> Result<String, PipelineError> {
98        if input.len() <= self.max_chars {
99            return Ok(input);
100        }
101        if self.max_chars == 0 {
102            return Ok(String::new());
103        }
104        // Find the last whitespace boundary within max_chars.
105        let truncated = &input[..self.max_chars];
106        let result = match truncated.rfind(|c: char| c.is_whitespace()) {
107            Some(pos) if pos > 0 => &truncated[..pos],
108            _ => truncated, // No word boundary found; hard-cut.
109        };
110        Ok(result.to_string())
111    }
112
113    fn name(&self) -> &str {
114        "TruncateStage"
115    }
116}
117
118/// Prepends a fixed string (e.g. a system-prompt prefix) to the input.
119pub struct PrependStage {
120    /// The string to prepend.
121    pub prefix: String,
122}
123
124#[async_trait]
125impl PipelineStage for PrependStage {
126    async fn process(&self, input: String) -> Result<String, PipelineError> {
127        Ok(format!("{}{}", self.prefix, input))
128    }
129
130    fn name(&self) -> &str {
131        "PrependStage"
132    }
133}
134
135/// Appends a fixed string (e.g. a context suffix or citation block) to the input.
136pub struct AppendStage {
137    /// The string to append.
138    pub suffix: String,
139}
140
141#[async_trait]
142impl PipelineStage for AppendStage {
143    async fn process(&self, input: String) -> Result<String, PipelineError> {
144        Ok(format!("{}{}", input, self.suffix))
145    }
146
147    fn name(&self) -> &str {
148        "AppendStage"
149    }
150}
151
152/// Performs a regex substitution on the input.
153///
154/// The `pattern` is compiled once at construction time and reused for every
155/// call to [`process`](PipelineStage::process).
156#[derive(Debug)]
157pub struct RegexReplaceStage {
158    /// The compiled regex pattern.
159    compiled: Regex,
160    /// Original pattern string, kept for error messages.
161    pattern_str: String,
162    /// Replacement string (supports `$1`, `${name}` capture group references).
163    pub replacement: String,
164}
165
166impl RegexReplaceStage {
167    /// Construct a new stage, compiling `pattern` immediately.
168    ///
169    /// Returns [`PipelineError::InvalidRegex`] if the pattern is invalid.
170    pub fn new(pattern: &str, replacement: &str) -> Result<Self, PipelineError> {
171        let compiled = Regex::new(pattern).map_err(|e| PipelineError::InvalidRegex {
172            pattern: pattern.to_string(),
173            source: e,
174        })?;
175        Ok(Self {
176            compiled,
177            pattern_str: pattern.to_string(),
178            replacement: replacement.to_string(),
179        })
180    }
181
182    /// Return the original pattern string.
183    pub fn pattern(&self) -> &str {
184        &self.pattern_str
185    }
186}
187
188#[async_trait]
189impl PipelineStage for RegexReplaceStage {
190    async fn process(&self, input: String) -> Result<String, PipelineError> {
191        Ok(self.compiled.replace_all(&input, self.replacement.as_str()).to_string())
192    }
193
194    fn name(&self) -> &str {
195        "RegexReplaceStage"
196    }
197}
198
199/// Detects whether the text is predominantly ASCII/Latin script or uses other
200/// Unicode scripts, then prepends a `[lang:en]` or `[lang:other]` tag.
201///
202/// The heuristic counts bytes: if ≥ 80 % of the non-whitespace characters are
203/// in the ASCII range (0x00–0x7F) the text is labelled `en`, otherwise `other`.
204pub struct LanguageDetectStage;
205
206#[async_trait]
207impl PipelineStage for LanguageDetectStage {
208    async fn process(&self, input: String) -> Result<String, PipelineError> {
209        let non_ws: Vec<char> = input.chars().filter(|c| !c.is_whitespace()).collect();
210        let tag = if non_ws.is_empty() {
211            "en"
212        } else {
213            let ascii_count = non_ws.iter().filter(|c| c.is_ascii()).count();
214            let ratio = ascii_count as f64 / non_ws.len() as f64;
215            if ratio >= 0.80 { "en" } else { "other" }
216        };
217        Ok(format!("[lang:{}] {}", tag, input))
218    }
219
220    fn name(&self) -> &str {
221        "LanguageDetectStage"
222    }
223}
224
225// ── Pipeline stats ─────────────────────────────────────────────────────────────
226
227/// Timing and length statistics for a single [`Pipeline::run`] call.
228#[derive(Debug, Clone, PartialEq)]
229pub struct PipelineStats {
230    /// Number of stages that were executed (including those that were no-ops).
231    pub stages_run: usize,
232    /// Length of the original input string in bytes.
233    pub input_len: usize,
234    /// Length of the final output string in bytes.
235    pub output_len: usize,
236    /// Wall-clock time taken by the entire pipeline in milliseconds.
237    pub elapsed_ms: u64,
238}
239
240// ── Pipeline run result ────────────────────────────────────────────────────────
241
242/// The result of a successful [`Pipeline::run`] call.
243#[derive(Debug, Clone)]
244pub struct PipelineResult {
245    /// The final transformed string.
246    pub output: String,
247    /// Execution statistics.
248    pub stats: PipelineStats,
249}
250
251// ── Pipeline ──────────────────────────────────────────────────────────────────
252
253/// An ordered sequence of [`PipelineStage`]s.
254///
255/// Each stage's output is fed as the next stage's input.  Build pipelines with
256/// [`PipelineBuilder`].
257pub struct Pipeline {
258    stages: Vec<Box<dyn PipelineStage>>,
259}
260
261impl Pipeline {
262    /// Execute all stages in order, returning the transformed string and
263    /// execution statistics.
264    ///
265    /// Fails fast: if any stage returns an error the pipeline stops and
266    /// propagates the error.
267    pub async fn run(&self, input: String) -> Result<PipelineResult, PipelineError> {
268        let start = Instant::now();
269        let input_len = input.len();
270        let mut current = input;
271
272        for stage in &self.stages {
273            current = stage.process(current).await?;
274        }
275
276        let elapsed_ms = start.elapsed().as_millis() as u64;
277        let output_len = current.len();
278        Ok(PipelineResult {
279            output: current,
280            stats: PipelineStats {
281                stages_run: self.stages.len(),
282                input_len,
283                output_len,
284                elapsed_ms,
285            },
286        })
287    }
288
289    /// Return the number of stages in this pipeline.
290    pub fn stage_count(&self) -> usize {
291        self.stages.len()
292    }
293}
294
295// ── Builder ───────────────────────────────────────────────────────────────────
296
297/// Fluent builder for constructing a [`Pipeline`].
298///
299/// ## Example
300///
301/// ```rust
302/// use tokio_prompt_orchestrator::pipeline::{PipelineBuilder, TrimStage, PrependStage};
303///
304/// let pipeline = PipelineBuilder::new()
305///     .add(TrimStage)
306///     .add(PrependStage { prefix: "System: ".to_string() })
307///     .build();
308/// ```
309#[derive(Default)]
310pub struct PipelineBuilder {
311    stages: Vec<Box<dyn PipelineStage>>,
312}
313
314impl PipelineBuilder {
315    /// Create an empty builder.
316    pub fn new() -> Self {
317        Self { stages: Vec::new() }
318    }
319
320    /// Append a stage to the pipeline.
321    #[allow(clippy::should_implement_trait)]
322    pub fn add<S: PipelineStage + 'static>(mut self, stage: S) -> Self {
323        self.stages.push(Box::new(stage));
324        self
325    }
326
327    /// Consume the builder and return the configured [`Pipeline`].
328    pub fn build(self) -> Pipeline {
329        Pipeline { stages: self.stages }
330    }
331}
332
333// ── Tests ──────────────────────────────────────────────────────────────────────
334
335#[cfg(test)]
336mod tests {
337    use super::*;
338
339    // ── TrimStage ─────────────────────────────────────────────────────────────
340
341    #[tokio::test]
342    async fn trim_stage_removes_leading_whitespace() {
343        let s = TrimStage;
344        let out = s.process("   hello".to_string()).await.unwrap();
345        assert_eq!(out, "hello");
346    }
347
348    #[tokio::test]
349    async fn trim_stage_removes_trailing_whitespace() {
350        let s = TrimStage;
351        let out = s.process("hello   ".to_string()).await.unwrap();
352        assert_eq!(out, "hello");
353    }
354
355    #[tokio::test]
356    async fn trim_stage_empty_string() {
357        let s = TrimStage;
358        let out = s.process("   ".to_string()).await.unwrap();
359        assert_eq!(out, "");
360    }
361
362    #[tokio::test]
363    async fn trim_stage_no_op_when_already_trimmed() {
364        let s = TrimStage;
365        let out = s.process("hello world".to_string()).await.unwrap();
366        assert_eq!(out, "hello world");
367    }
368
369    // ── TruncateStage ─────────────────────────────────────────────────────────
370
371    #[tokio::test]
372    async fn truncate_stage_short_input_unchanged() {
373        let s = TruncateStage { max_chars: 100 };
374        let out = s.process("hello".to_string()).await.unwrap();
375        assert_eq!(out, "hello");
376    }
377
378    #[tokio::test]
379    async fn truncate_stage_breaks_at_word_boundary() {
380        let s = TruncateStage { max_chars: 10 };
381        let out = s.process("hello world foo".to_string()).await.unwrap();
382        // "hello worl" -> last space at index 5 -> "hello"
383        assert_eq!(out, "hello");
384    }
385
386    #[tokio::test]
387    async fn truncate_stage_zero_max_chars() {
388        let s = TruncateStage { max_chars: 0 };
389        let out = s.process("hello".to_string()).await.unwrap();
390        assert_eq!(out, "");
391    }
392
393    #[tokio::test]
394    async fn truncate_stage_exact_length() {
395        let s = TruncateStage { max_chars: 5 };
396        let out = s.process("hello".to_string()).await.unwrap();
397        assert_eq!(out, "hello");
398    }
399
400    #[tokio::test]
401    async fn truncate_stage_no_word_boundary_hard_cuts() {
402        let s = TruncateStage { max_chars: 3 };
403        let out = s.process("hello".to_string()).await.unwrap();
404        assert_eq!(out, "hel");
405    }
406
407    // ── PrependStage ──────────────────────────────────────────────────────────
408
409    #[tokio::test]
410    async fn prepend_stage_adds_prefix() {
411        let s = PrependStage { prefix: "System: ".to_string() };
412        let out = s.process("hello".to_string()).await.unwrap();
413        assert_eq!(out, "System: hello");
414    }
415
416    #[tokio::test]
417    async fn prepend_stage_empty_prefix() {
418        let s = PrependStage { prefix: String::new() };
419        let out = s.process("hello".to_string()).await.unwrap();
420        assert_eq!(out, "hello");
421    }
422
423    // ── AppendStage ───────────────────────────────────────────────────────────
424
425    #[tokio::test]
426    async fn append_stage_adds_suffix() {
427        let s = AppendStage { suffix: " [END]".to_string() };
428        let out = s.process("hello".to_string()).await.unwrap();
429        assert_eq!(out, "hello [END]");
430    }
431
432    #[tokio::test]
433    async fn append_stage_empty_suffix() {
434        let s = AppendStage { suffix: String::new() };
435        let out = s.process("hello".to_string()).await.unwrap();
436        assert_eq!(out, "hello");
437    }
438
439    // ── RegexReplaceStage ─────────────────────────────────────────────────────
440
441    #[tokio::test]
442    async fn regex_replace_stage_basic_substitution() {
443        let s = RegexReplaceStage::new(r"\bfoo\b", "bar").unwrap();
444        let out = s.process("foo is foo".to_string()).await.unwrap();
445        assert_eq!(out, "bar is bar");
446    }
447
448    #[tokio::test]
449    async fn regex_replace_stage_no_match_unchanged() {
450        let s = RegexReplaceStage::new(r"\bxyz\b", "abc").unwrap();
451        let out = s.process("hello world".to_string()).await.unwrap();
452        assert_eq!(out, "hello world");
453    }
454
455    #[tokio::test]
456    async fn regex_replace_stage_invalid_pattern_errors() {
457        let result = RegexReplaceStage::new(r"[invalid", "x");
458        assert!(result.is_err());
459        assert!(matches!(result.unwrap_err(), PipelineError::InvalidRegex { .. }));
460    }
461
462    #[tokio::test]
463    async fn regex_replace_stage_capture_group_replacement() {
464        let s = RegexReplaceStage::new(r"(\w+)\s+(\w+)", "$2 $1").unwrap();
465        let out = s.process("hello world".to_string()).await.unwrap();
466        assert_eq!(out, "world hello");
467    }
468
469    // ── LanguageDetectStage ───────────────────────────────────────────────────
470
471    #[tokio::test]
472    async fn language_detect_ascii_text_tagged_en() {
473        let s = LanguageDetectStage;
474        let out = s.process("Hello world".to_string()).await.unwrap();
475        assert!(out.starts_with("[lang:en]"));
476    }
477
478    #[tokio::test]
479    async fn language_detect_other_script_tagged_other() {
480        let s = LanguageDetectStage;
481        // Cyrillic text
482        let out = s.process("Привет мир".to_string()).await.unwrap();
483        assert!(out.starts_with("[lang:other]"));
484    }
485
486    #[tokio::test]
487    async fn language_detect_empty_string_tagged_en() {
488        let s = LanguageDetectStage;
489        let out = s.process(String::new()).await.unwrap();
490        assert!(out.starts_with("[lang:en]"));
491    }
492
493    // ── Pipeline ──────────────────────────────────────────────────────────────
494
495    #[tokio::test]
496    async fn pipeline_runs_stages_in_order() {
497        let pipeline = PipelineBuilder::new()
498            .add(TrimStage)
499            .add(PrependStage { prefix: ">>".to_string() })
500            .build();
501
502        let result = pipeline.run("  hello  ".to_string()).await.unwrap();
503        assert_eq!(result.output, ">>hello");
504    }
505
506    #[tokio::test]
507    async fn pipeline_stats_stage_count() {
508        let pipeline = PipelineBuilder::new()
509            .add(TrimStage)
510            .add(AppendStage { suffix: "!".to_string() })
511            .build();
512
513        let result = pipeline.run("hi".to_string()).await.unwrap();
514        assert_eq!(result.stats.stages_run, 2);
515    }
516
517    #[tokio::test]
518    async fn pipeline_stats_input_output_len() {
519        let pipeline = PipelineBuilder::new()
520            .add(AppendStage { suffix: "XYZ".to_string() })
521            .build();
522
523        let result = pipeline.run("hello".to_string()).await.unwrap();
524        assert_eq!(result.stats.input_len, 5);
525        assert_eq!(result.stats.output_len, 8);
526    }
527
528    #[tokio::test]
529    async fn pipeline_empty_no_stages() {
530        let pipeline = PipelineBuilder::new().build();
531        let result = pipeline.run("hello".to_string()).await.unwrap();
532        assert_eq!(result.output, "hello");
533        assert_eq!(result.stats.stages_run, 0);
534    }
535
536    #[tokio::test]
537    async fn pipeline_full_chain() {
538        let regex_stage = RegexReplaceStage::new(r"\bLLM\b", "AI model").unwrap();
539        let pipeline = PipelineBuilder::new()
540            .add(TrimStage)
541            .add(TruncateStage { max_chars: 200 })
542            .add(PrependStage { prefix: "Context: ".to_string() })
543            .add(AppendStage { suffix: " [end]".to_string() })
544            .add(regex_stage)
545            .build();
546
547        let input = "  Ask the LLM to summarize.  ";
548        let result = pipeline.run(input.to_string()).await.unwrap();
549        assert!(result.output.contains("AI model"));
550        assert!(result.output.starts_with("Context: "));
551        assert!(result.output.ends_with("[end]"));
552    }
553
554    #[tokio::test]
555    async fn pipeline_builder_stage_count() {
556        let pipeline = PipelineBuilder::new()
557            .add(TrimStage)
558            .add(TrimStage)
559            .add(TrimStage)
560            .build();
561        assert_eq!(pipeline.stage_count(), 3);
562    }
563}