tokio_prompt_orchestrator/
multi_pipeline.rs1use 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#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
65pub enum PromptClass {
66 Faq,
69 Reasoning,
72 Code,
75 Document,
78 Batch,
80 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
98pub struct HeuristicClassifier {
104 pub document_token_threshold: usize,
106 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 fn approx_tokens(text: &str) -> usize {
122 text.split_whitespace().count()
123 }
124
125 pub fn classify(&self, prompt: &str) -> PromptClass {
129 let lower = prompt.to_lowercase();
130 let tokens = Self::approx_tokens(prompt);
131
132 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 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 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 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 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
197pub struct PipelineDescriptor {
199 pub name: String,
201 pub target_class: PromptClass,
203 pub also_serves: Vec<PromptClass>,
205 pub worker: Arc<dyn ModelWorker>,
207 pub config: Option<crate::config::PipelineConfig>,
209}
210
211impl PipelineDescriptor {
212 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 pub fn also_serving(mut self, classes: Vec<PromptClass>) -> Self {
229 self.also_serves = classes;
230 self
231 }
232
233 pub fn with_config(mut self, config: crate::config::PipelineConfig) -> Self {
235 self.config = Some(config);
236 self
237 }
238}
239
240#[derive(Default)]
242struct PipelineStats {
243 routed: AtomicU64,
244 shed: AtomicU64,
245 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
275struct LivePipeline {
277 descriptor_name: String,
278 target_class: PromptClass,
279 also_serves: Vec<PromptClass>,
280 handles: PipelineHandles,
281 stats: Arc<PipelineStats>,
282}
283
284pub struct MultiPipelineRouter {
288 pipelines: Vec<LivePipeline>,
289 classifier: HeuristicClassifier,
290 stats_map: Arc<DashMap<String, Arc<PipelineStats>>>,
292}
293
294impl MultiPipelineRouter {
295 pub fn builder() -> MultiPipelineRouterBuilder {
297 MultiPipelineRouterBuilder::new()
298 }
299
300 pub fn classify(&self, request: &PromptRequest) -> PromptClass {
305 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 pub async fn route(&self, request: PromptRequest) -> Result<(), OrchestratorError> {
331 let class = self.classify(&request);
332
333 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 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 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#[derive(Debug, Clone, Serialize)]
399pub struct PipelineRoutingStats {
400 pub name: String,
402 pub target_class: PromptClass,
404 pub routed: u64,
406 pub shed: u64,
408 pub ema_latency_ms: f64,
410}
411
412pub struct MultiPipelineRouterBuilder {
414 descriptors: Vec<PipelineDescriptor>,
415}
416
417impl MultiPipelineRouterBuilder {
418 fn new() -> Self {
419 Self {
420 descriptors: Vec::new(),
421 }
422 }
423
424 pub fn add_pipeline(mut self, descriptor: PipelineDescriptor) -> Self {
426 self.descriptors.push(descriptor);
427 self
428 }
429
430 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 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 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 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 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}