Skip to main content

tokio_prompt_orchestrator/
multi_pipeline.rs

1//! Multi-pipeline routing with prompt classification.
2//!
3//! Manages a fleet of named pipeline instances and routes each incoming
4//! [`PromptRequest`] to the best-fit pipeline based on prompt characteristics:
5//! length, complexity signals, explicit priority hints, and learned latency.
6//!
7//! ## Concept
8//!
9//! A single LLM orchestrator often needs to handle very different workloads
10//! simultaneously:
11//!
12//! - **Fast path** — short FAQ-style questions, single-turn, low latency required
13//! - **Reasoning path** — multi-step problems, chain-of-thought, higher quality
14//! - **Code path** — specialised coding model, longer context window
15//! - **Batch path** — offline background jobs, throughput over latency
16//!
17//! Rather than forcing all traffic through one pipeline, `MultiPipelineRouter`
18//! dispatches to the right pipeline and falls back gracefully when a pipeline
19//! is at capacity.
20//!
21//! ## Example
22//!
23//! ```no_run
24//! use std::sync::Arc;
25//! use tokio_prompt_orchestrator::{EchoWorker, PromptRequest, SessionId};
26//! use tokio_prompt_orchestrator::multi_pipeline::{
27//!     MultiPipelineRouter, PipelineDescriptor, PromptClass,
28//! };
29//! use std::collections::HashMap;
30//!
31//! # async fn example() {
32//! let router = MultiPipelineRouter::builder()
33//!     .add_pipeline(PipelineDescriptor::new("fast", PromptClass::Faq, Arc::new(EchoWorker::new())))
34//!     .add_pipeline(PipelineDescriptor::new("reasoning", PromptClass::Reasoning, Arc::new(EchoWorker::new())))
35//!     .build();
36//!
37//! let request = PromptRequest {
38//!     session: SessionId::new("s1"),
39//!     request_id: "r1".to_string(),
40//!     input: "What is 2+2?".to_string(),
41//!     meta: HashMap::new(),
42//!     deadline: None,
43//! };
44//!
45//! let class = router.classify(&request);
46//! router.route(request).await.unwrap();
47//! # }
48//! ```
49
50use crate::{
51    stages::{spawn_pipeline_with_config, PipelineHandles},
52    worker::ModelWorker,
53    OrchestratorError, PromptRequest,
54};
55use dashmap::DashMap;
56use serde::{Deserialize, Serialize};
57use std::sync::{
58    atomic::{AtomicU64, Ordering},
59    Arc,
60};
61use tracing::{debug, info, warn};
62
63/// Broad classification of a prompt's intent and resource requirements.
64#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
65pub enum PromptClass {
66    /// Short, simple questions with a known answer pattern.
67    /// Routes to the fast pipeline for minimal latency.
68    Faq,
69    /// Multi-step reasoning, analysis, or planning tasks.
70    /// Routes to the reasoning pipeline for quality.
71    Reasoning,
72    /// Code generation, review, or debugging requests.
73    /// Routes to the code-specialised pipeline.
74    Code,
75    /// Long-form document processing (summarisation, translation, extraction).
76    /// Routes to the batch pipeline for throughput.
77    Document,
78    /// Background / offline work — latency-insensitive.
79    Batch,
80    /// Catch-all: no strong classification signal.
81    General,
82}
83
84impl std::fmt::Display for PromptClass {
85    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
86        let s = match self {
87            Self::Faq => "faq",
88            Self::Reasoning => "reasoning",
89            Self::Code => "code",
90            Self::Document => "document",
91            Self::Batch => "batch",
92            Self::General => "general",
93        };
94        write!(f, "{s}")
95    }
96}
97
98/// Heuristic prompt classifier.
99///
100/// Uses lightweight, allocation-free checks to classify a prompt without
101/// calling any external model. For production use, swap in an embedding-based
102/// classifier via `PromptClassifier` trait implementations.
103pub struct HeuristicClassifier {
104    /// Minimum token-estimate for a prompt to be treated as a Document.
105    pub document_token_threshold: usize,
106    /// Minimum token-estimate for reasoning (above FAQ, below document).
107    pub reasoning_token_threshold: usize,
108}
109
110impl Default for HeuristicClassifier {
111    fn default() -> Self {
112        Self {
113            document_token_threshold: 800,
114            reasoning_token_threshold: 100,
115        }
116    }
117}
118
119impl HeuristicClassifier {
120    /// Quick token count approximation: split on whitespace.
121    fn approx_tokens(text: &str) -> usize {
122        text.split_whitespace().count()
123    }
124
125    /// Classify a prompt text into a [`PromptClass`].
126    ///
127    /// Classification is purely heuristic and fast (~50 ns).
128    pub fn classify(&self, prompt: &str) -> PromptClass {
129        let lower = prompt.to_lowercase();
130        let tokens = Self::approx_tokens(prompt);
131
132        // Code signals
133        if lower.contains("```")
134            || lower.contains("fn ")
135            || lower.contains("def ")
136            || lower.contains("class ")
137            || lower.contains("impl ")
138            || lower.contains("function ")
139            || lower.contains("// ")
140            || lower.contains("# code")
141            || lower.contains("write code")
142            || lower.contains("debug")
143            || lower.contains("compile error")
144        {
145            return PromptClass::Code;
146        }
147
148        // Reasoning signals
149        if lower.contains("step by step")
150            || lower.contains("explain why")
151            || lower.contains("analyse")
152            || lower.contains("analyze")
153            || lower.contains("compare")
154            || lower.contains("pros and cons")
155            || lower.contains("reason")
156            || lower.contains("chain of thought")
157            || lower.contains("think through")
158        {
159            return PromptClass::Reasoning;
160        }
161
162        // Document signals
163        if tokens >= self.document_token_threshold
164            || lower.contains("summarise")
165            || lower.contains("summarize")
166            || lower.contains("translate")
167            || lower.contains("extract all")
168            || lower.contains("the following document")
169            || lower.contains("the following text")
170        {
171            return PromptClass::Document;
172        }
173
174        // Batch signals (explicit meta hints checked in router)
175        if lower.contains("batch")
176            || lower.contains("background")
177            || lower.contains("offline")
178            || lower.contains("no rush")
179        {
180            return PromptClass::Batch;
181        }
182
183        // FAQ: short and direct
184        if tokens < self.reasoning_token_threshold
185            && (lower.ends_with('?') || lower.starts_with("what") || lower.starts_with("who")
186                || lower.starts_with("when") || lower.starts_with("where")
187                || lower.starts_with("how many") || lower.starts_with("is ")
188                || lower.starts_with("does ") || lower.starts_with("can "))
189        {
190            return PromptClass::Faq;
191        }
192
193        PromptClass::General
194    }
195}
196
197/// Descriptor for a single named pipeline instance.
198pub struct PipelineDescriptor {
199    /// Unique name (e.g. "fast", "reasoning", "code").
200    pub name: String,
201    /// Primary prompt class this pipeline serves.
202    pub target_class: PromptClass,
203    /// Additional classes this pipeline can serve (fallback order).
204    pub also_serves: Vec<PromptClass>,
205    /// The worker for this pipeline.
206    pub worker: Arc<dyn ModelWorker>,
207    /// Pipeline config override (uses default if None).
208    pub config: Option<crate::config::PipelineConfig>,
209}
210
211impl PipelineDescriptor {
212    /// Create a basic descriptor serving one prompt class.
213    pub fn new(
214        name: impl Into<String>,
215        target_class: PromptClass,
216        worker: Arc<dyn ModelWorker>,
217    ) -> Self {
218        Self {
219            name: name.into(),
220            target_class,
221            also_serves: Vec::new(),
222            worker,
223            config: None,
224        }
225    }
226
227    /// Add additional prompt classes this pipeline can handle.
228    pub fn also_serving(mut self, classes: Vec<PromptClass>) -> Self {
229        self.also_serves = classes;
230        self
231    }
232
233    /// Override the default pipeline config.
234    pub fn with_config(mut self, config: crate::config::PipelineConfig) -> Self {
235        self.config = Some(config);
236        self
237    }
238}
239
240/// Per-pipeline runtime stats tracked for routing decisions.
241#[derive(Default)]
242struct PipelineStats {
243    routed: AtomicU64,
244    shed: AtomicU64,
245    /// EMA of recent latency in milliseconds (×1000 as integer for atomics).
246    ema_latency_ms_x1000: AtomicU64,
247}
248
249impl PipelineStats {
250    fn record_routed(&self) {
251        self.routed.fetch_add(1, Ordering::Relaxed);
252    }
253
254    fn record_shed(&self) {
255        self.shed.fetch_add(1, Ordering::Relaxed);
256    }
257
258    fn update_latency(&self, latency_ms: f64) {
259        const ALPHA: f64 = 0.1;
260        let current = self.ema_latency_ms_x1000.load(Ordering::Relaxed) as f64 / 1000.0;
261        let updated = if current == 0.0 {
262            latency_ms
263        } else {
264            ALPHA * latency_ms + (1.0 - ALPHA) * current
265        };
266        self.ema_latency_ms_x1000
267            .store((updated * 1000.0) as u64, Ordering::Relaxed);
268    }
269
270    fn ema_latency_ms(&self) -> f64 {
271        self.ema_latency_ms_x1000.load(Ordering::Relaxed) as f64 / 1000.0
272    }
273}
274
275/// A live, spawned pipeline entry.
276struct LivePipeline {
277    descriptor_name: String,
278    target_class: PromptClass,
279    also_serves: Vec<PromptClass>,
280    handles: PipelineHandles,
281    stats: Arc<PipelineStats>,
282}
283
284/// Routes incoming prompts across multiple named pipeline instances.
285///
286/// Build via [`MultiPipelineRouter::builder()`].
287pub struct MultiPipelineRouter {
288    pipelines: Vec<LivePipeline>,
289    classifier: HeuristicClassifier,
290    /// Pipeline name → stats (for metrics/dashboard).
291    stats_map: Arc<DashMap<String, Arc<PipelineStats>>>,
292}
293
294impl MultiPipelineRouter {
295    /// Create a new builder.
296    pub fn builder() -> MultiPipelineRouterBuilder {
297        MultiPipelineRouterBuilder::new()
298    }
299
300    /// Classify a prompt request using the heuristic classifier.
301    ///
302    /// Respects an optional `"pipeline_class"` key in `request.meta` for
303    /// explicit override.
304    pub fn classify(&self, request: &PromptRequest) -> PromptClass {
305        // Check for explicit override in meta
306        if let Some(class_hint) = request.meta.get("pipeline_class") {
307            match class_hint.as_str() {
308                "faq" => return PromptClass::Faq,
309                "reasoning" => return PromptClass::Reasoning,
310                "code" => return PromptClass::Code,
311                "document" => return PromptClass::Document,
312                "batch" => return PromptClass::Batch,
313                _ => {}
314            }
315        }
316        self.classifier.classify(&request.input)
317    }
318
319    /// Route a request to the best-fit pipeline.
320    ///
321    /// Routing priority:
322    /// 1. Pipeline whose `target_class` matches the classified prompt.
323    /// 2. Any pipeline that lists the class in `also_serves`.
324    /// 3. First pipeline (default/general fallback).
325    ///
326    /// # Errors
327    ///
328    /// Returns [`OrchestratorError::ChannelClosed`] if the selected
329    /// pipeline's input channel is closed.
330    pub async fn route(&self, request: PromptRequest) -> Result<(), OrchestratorError> {
331        let class = self.classify(&request);
332
333        // Find primary match
334        let pipeline = self
335            .pipelines
336            .iter()
337            .find(|p| p.target_class == class)
338            .or_else(|| {
339                self.pipelines
340                    .iter()
341                    .find(|p| p.also_serves.contains(&class))
342            })
343            .or_else(|| self.pipelines.first());
344
345        let Some(pipeline) = pipeline else {
346            return Err(OrchestratorError::ConfigError(
347                "MultiPipelineRouter has no pipelines registered".to_string(),
348            ));
349        };
350
351        debug!(
352            pipeline = %pipeline.descriptor_name,
353            class = %class,
354            request_id = %request.request_id,
355            "routing request"
356        );
357
358        pipeline.stats.record_routed();
359
360        pipeline
361            .handles
362            .input_tx
363            .send(request)
364            .await
365            .map_err(|_| {
366                pipeline.stats.record_shed();
367                warn!(
368                    pipeline = %pipeline.descriptor_name,
369                    "pipeline channel closed — request shed"
370                );
371                OrchestratorError::ChannelClosed
372            })
373    }
374
375    /// Return a snapshot of routing statistics per pipeline.
376    pub fn stats(&self) -> Vec<PipelineRoutingStats> {
377        self.pipelines
378            .iter()
379            .map(|p| PipelineRoutingStats {
380                name: p.descriptor_name.clone(),
381                target_class: p.target_class,
382                routed: p.stats.routed.load(Ordering::Relaxed),
383                shed: p.stats.shed.load(Ordering::Relaxed),
384                ema_latency_ms: p.stats.ema_latency_ms(),
385            })
386            .collect()
387    }
388
389    /// Record observed latency for a pipeline (called by the pipeline consumer).
390    pub fn record_latency(&self, pipeline_name: &str, latency_ms: f64) {
391        if let Some(stats) = self.stats_map.get(pipeline_name) {
392            stats.update_latency(latency_ms);
393        }
394    }
395}
396
397/// Snapshot of routing stats for one pipeline.
398#[derive(Debug, Clone, Serialize)]
399pub struct PipelineRoutingStats {
400    /// Pipeline name.
401    pub name: String,
402    /// Primary class this pipeline targets.
403    pub target_class: PromptClass,
404    /// Total requests routed to this pipeline.
405    pub routed: u64,
406    /// Requests shed (channel full or closed).
407    pub shed: u64,
408    /// Exponential moving average of observed latency.
409    pub ema_latency_ms: f64,
410}
411
412/// Builder for [`MultiPipelineRouter`].
413pub struct MultiPipelineRouterBuilder {
414    descriptors: Vec<PipelineDescriptor>,
415}
416
417impl MultiPipelineRouterBuilder {
418    fn new() -> Self {
419        Self {
420            descriptors: Vec::new(),
421        }
422    }
423
424    /// Register a pipeline.
425    pub fn add_pipeline(mut self, descriptor: PipelineDescriptor) -> Self {
426        self.descriptors.push(descriptor);
427        self
428    }
429
430    /// Spawn all pipelines and return the router.
431    pub fn build(self) -> MultiPipelineRouter {
432        let stats_map: Arc<DashMap<String, Arc<PipelineStats>>> = Arc::new(DashMap::new());
433        let mut pipelines = Vec::new();
434
435        for desc in self.descriptors {
436            let stats = Arc::new(PipelineStats::default());
437            stats_map.insert(desc.name.clone(), Arc::clone(&stats));
438
439            let handles = if let Some(ref config) = desc.config {
440                spawn_pipeline_with_config(desc.worker, config)
441            } else {
442                crate::stages::spawn_pipeline(desc.worker)
443            };
444
445            info!(
446                pipeline = %desc.name,
447                class = %desc.target_class,
448                "pipeline spawned in multi-pipeline router"
449            );
450
451            pipelines.push(LivePipeline {
452                descriptor_name: desc.name,
453                target_class: desc.target_class,
454                also_serves: desc.also_serves,
455                handles,
456                stats,
457            });
458        }
459
460        MultiPipelineRouter {
461            pipelines,
462            classifier: HeuristicClassifier::default(),
463            stats_map,
464        }
465    }
466}
467
468#[cfg(test)]
469mod tests {
470    use super::*;
471
472    #[test]
473    fn test_classify_code() {
474        let c = HeuristicClassifier::default();
475        assert_eq!(c.classify("```rust\nfn main() {}\n```"), PromptClass::Code);
476        assert_eq!(c.classify("debug this compile error"), PromptClass::Code);
477        assert_eq!(c.classify("write code to sort a list"), PromptClass::Code);
478    }
479
480    #[test]
481    fn test_classify_faq() {
482        let c = HeuristicClassifier::default();
483        assert_eq!(c.classify("What is Rust?"), PromptClass::Faq);
484        assert_eq!(c.classify("Who invented TCP/IP?"), PromptClass::Faq);
485        assert_eq!(c.classify("Is Tokio async?"), PromptClass::Faq);
486    }
487
488    #[test]
489    fn test_classify_reasoning() {
490        let c = HeuristicClassifier::default();
491        assert_eq!(
492            c.classify("Explain why Rust is memory safe step by step"),
493            PromptClass::Reasoning
494        );
495        assert_eq!(
496            c.classify("analyse the pros and cons of async Rust"),
497            PromptClass::Reasoning
498        );
499    }
500
501    #[test]
502    fn test_classify_document() {
503        let c = HeuristicClassifier::default();
504        // Long text crosses document threshold
505        let long = "word ".repeat(900);
506        assert_eq!(c.classify(&long), PromptClass::Document);
507        assert_eq!(c.classify("summarize the following document"), PromptClass::Document);
508    }
509
510    #[test]
511    fn test_classify_general_fallback() {
512        let c = HeuristicClassifier::default();
513        // Non-signal text
514        assert_eq!(c.classify("hello"), PromptClass::General);
515    }
516
517    #[test]
518    fn test_pipeline_stats_ema() {
519        let stats = PipelineStats::default();
520        stats.update_latency(100.0);
521        stats.update_latency(200.0);
522        let ema = stats.ema_latency_ms();
523        // Should be between 100 and 200
524        assert!(ema > 100.0 && ema < 200.0, "ema={ema}");
525    }
526
527    #[tokio::test]
528    async fn test_router_routes_to_matching_pipeline() {
529        use crate::EchoWorker;
530        use std::collections::HashMap;
531
532        let router = MultiPipelineRouter::builder()
533            .add_pipeline(PipelineDescriptor::new(
534                "faq",
535                PromptClass::Faq,
536                Arc::new(EchoWorker::new()),
537            ))
538            .add_pipeline(PipelineDescriptor::new(
539                "general",
540                PromptClass::General,
541                Arc::new(EchoWorker::new()),
542            ))
543            .build();
544
545        let req = PromptRequest {
546            session: crate::SessionId::new("s1"),
547            request_id: "r1".to_string(),
548            input: "What is Tokio?".to_string(),
549            meta: HashMap::new(),
550            deadline: None,
551        };
552
553        let class = router.classify(&req);
554        assert_eq!(class, PromptClass::Faq);
555
556        // Route should succeed without error
557        router.route(req).await.expect("route should succeed");
558
559        let stats = router.stats();
560        let faq_stats = stats.iter().find(|s| s.name == "faq").unwrap();
561        assert_eq!(faq_stats.routed, 1);
562    }
563
564    #[tokio::test]
565    async fn test_meta_override() {
566        use crate::EchoWorker;
567        use std::collections::HashMap;
568
569        let router = MultiPipelineRouter::builder()
570            .add_pipeline(PipelineDescriptor::new(
571                "general",
572                PromptClass::General,
573                Arc::new(EchoWorker::new()),
574            ))
575            .build();
576
577        let mut meta = HashMap::new();
578        meta.insert("pipeline_class".to_string(), "code".to_string());
579        let req = PromptRequest {
580            session: crate::SessionId::new("s1"),
581            request_id: "r1".to_string(),
582            input: "hello world".to_string(),
583            meta,
584            deadline: None,
585        };
586
587        assert_eq!(router.classify(&req), PromptClass::Code);
588    }
589}