tokio_prompt_orchestrator/
pipeline.rs1use async_trait::async_trait;
25use regex::Regex;
26use std::time::Instant;
27use thiserror::Error;
28
29#[derive(Debug, Error)]
33pub enum PipelineError {
34 #[error("stage '{stage}' failed: {message}")]
36 StageError {
37 stage: String,
39 message: String,
41 },
42
43 #[error("invalid regex pattern '{pattern}': {source}")]
45 InvalidRegex {
46 pattern: String,
48 #[source]
50 source: regex::Error,
51 },
52}
53
54#[async_trait]
62pub trait PipelineStage: Send + Sync {
63 async fn process(&self, input: String) -> Result<String, PipelineError>;
65
66 fn name(&self) -> &str;
68}
69
70pub 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
86pub struct TruncateStage {
91 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 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, };
110 Ok(result.to_string())
111 }
112
113 fn name(&self) -> &str {
114 "TruncateStage"
115 }
116}
117
118pub struct PrependStage {
120 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
135pub struct AppendStage {
137 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#[derive(Debug)]
157pub struct RegexReplaceStage {
158 compiled: Regex,
160 pattern_str: String,
162 pub replacement: String,
164}
165
166impl RegexReplaceStage {
167 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 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
199pub 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#[derive(Debug, Clone, PartialEq)]
229pub struct PipelineStats {
230 pub stages_run: usize,
232 pub input_len: usize,
234 pub output_len: usize,
236 pub elapsed_ms: u64,
238}
239
240#[derive(Debug, Clone)]
244pub struct PipelineResult {
245 pub output: String,
247 pub stats: PipelineStats,
249}
250
251pub struct Pipeline {
258 stages: Vec<Box<dyn PipelineStage>>,
259}
260
261impl Pipeline {
262 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 pub fn stage_count(&self) -> usize {
291 self.stages.len()
292 }
293}
294
295#[derive(Default)]
310pub struct PipelineBuilder {
311 stages: Vec<Box<dyn PipelineStage>>,
312}
313
314impl PipelineBuilder {
315 pub fn new() -> Self {
317 Self { stages: Vec::new() }
318 }
319
320 #[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 pub fn build(self) -> Pipeline {
329 Pipeline { stages: self.stages }
330 }
331}
332
333#[cfg(test)]
336mod tests {
337 use super::*;
338
339 #[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 #[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 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 #[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 #[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 #[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 #[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 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 #[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}