tokio_prompt_orchestrator/lib.rs
1//! Put a queue, deduplication and a circuit breaker between your app and an LLM
2//! provider, so repeated prompts cost one call and an outage fails fast instead
3//! of piling up.
4//!
5//! 
6//!
7//! ## Quick start
8//!
9//! ```toml
10//! [dependencies]
11//! tokio-prompt-orchestrator = "1.4"
12//! tokio = { version = "1", features = ["rt-multi-thread", "macros"] }
13//! ```
14//!
15//! Send one prompt through the pipeline with the offline [`EchoWorker`] (no API
16//! key), read the answer, and check the dead-letter queue for anything dropped:
17//!
18//! ```
19//! use std::{collections::HashMap, sync::Arc};
20//! use tokio_prompt_orchestrator::{spawn_pipeline, EchoWorker, ModelWorker, PromptRequest, SessionId};
21//!
22//! #[tokio::main]
23//! async fn main() -> Result<(), Box<dyn std::error::Error>> {
24//! // Swap EchoWorker for AnthropicWorker, OpenAiWorker, LlamaCppWorker or VllmWorker.
25//! let worker: Arc<dyn ModelWorker> = Arc::new(EchoWorker::new());
26//! let handles = spawn_pipeline(worker);
27//! let mut output = handles.take_output_rx().await.ok_or("output already taken")?;
28//!
29//! handles.input_tx.send(PromptRequest {
30//! session: SessionId::new("demo"),
31//! request_id: "req-1".into(),
32//! input: "Hello, pipeline!".into(),
33//! meta: HashMap::new(),
34//! deadline: None,
35//! }).await?;
36//!
37//! let answer = output.recv().await.ok_or("pipeline closed")?;
38//! assert!(answer.text.contains("Hello, pipeline!"));
39//! println!("{}", answer.text);
40//!
41//! for dropped in handles.dlq.drain() {
42//! println!("dropped {}: {}", dropped.request_id, dropped.reason);
43//! }
44//! Ok(())
45//! }
46//! ```
47//!
48//! For a fuller example with deduplication and a simulated provider outage, run
49//! `cargo run --example llm_pipeline` in the repository.
50//!
51//! ## Main types
52//!
53//! - [`spawn_pipeline`] and [`spawn_pipeline_with_config`] start the five stages and return [`PipelineHandles`].
54//! - [`PromptRequest`] goes in through [`PipelineHandles::input_tx`]; [`PostOutput`] comes out of [`PipelineHandles::take_output_rx`].
55//! - [`ModelWorker`] is the trait for a backend: [`AnthropicWorker`], [`OpenAiWorker`], [`LlamaCppWorker`], [`VllmWorker`], [`EchoWorker`].
56//! - [`enhanced::Deduplicator`] and [`enhanced::CircuitBreaker`] are the resilience building blocks to wrap around a worker.
57//! - [`DeadLetterQueue`] holds every request that was dropped, with the reason.
58//! - [`PipelineConfig`] loads the same pipeline from TOML.
59//!
60//! ## How it works
61//!
62//! ```text
63//! PromptRequest -> RAG(512) -> Assemble(512) -> Inference(1024) -> Post(512) -> Stream(256)
64//! ```
65//!
66//! Each stage runs as its own Tokio task, joined by bounded channels. When a
67//! downstream channel is full, [`send_with_shed`] drops the item into the
68//! [`DeadLetterQueue`] instead of blocking the stage above it. Stage 3 checks the
69//! request deadline, calls the worker through a shared circuit breaker and
70//! enforces a timeout.
71//!
72//! ## Binaries and features
73//!
74//! The `orchestrator` binary (prebuilt on the
75//! [releases page](https://github.com/Mattbusel/tokio-prompt-orchestrator/releases/latest),
76//! or `cargo binstall tokio-prompt-orchestrator`) serves the pipeline over a
77//! terminal prompt and an HTTP API: `orchestrator --provider echo` runs with no key.
78//! No features are on by default. `web-api` adds the REST, SSE and WebSocket
79//! server, `full` adds metrics, Redis caching and rate limiting, `tui` a terminal
80//! dashboard, `mcp` an MCP server, and `self-improving` a self-tuning control loop.
81//!
82//! ## Modules
83//!
84//! | Module | Description |
85//! |--------|-------------|
86//! | [`stages`] | Five pipeline stage implementations and channel wiring |
87//! | [`worker`] | [`ModelWorker`] trait and the provider implementations |
88//! | [`enhanced`] | Resilience primitives: circuit breaker, dedup, semantic dedup (SimHash), retry, cache, rate limiter, smart batching |
89//! | [`config`] | TOML-deserialisable [`PipelineConfig`] with hot-reload support |
90//! | [`metrics`] | Prometheus metrics initialisation and helper functions |
91//! | [`routing`] | [`ModelRouter`] for complexity-scored routing; [`ArbitrageEngine`] for SLA-aware provider selection; [`PoolSizer`] for worker scaling |
92//! | [`security`] | [`PromptGuard`]: prompt injection and jailbreak detection (no external I/O) |
93//! | [`session`] | Multi-turn conversation context: injects history per session |
94//! | [`cascade`] | Multi-turn cascading inference: tool call loops with pluggable executors |
95//! | [`multi_pipeline`] | Named pipeline fleet with prompt classification and per-class routing |
96//! | [`ab_test`] | Prompt A/B testing: consistent hashing assignment, Welch's t-test, Cohen's d |
97//! | [`coordination`] | Agent fleet management and task claiming |
98//! | [`adaptive_pool`] | Kalman-filter worker pool controller: predicts queue depth and recommends scale events |
99//! | `web_api` | REST/SSE/WebSocket server (feature: `web-api`) |
100//! | `distributed` | Redis dedup and NATS coordination (feature: `distributed`) |
101//! | `tui` | Ratatui terminal dashboard (feature: `tui`) |
102//! | `self_tune`, `self_modify`, `intelligence`, `evolution`, `self_improve` | Self-improving control loop (features of the same names) |
103//!
104//! [`PipelineConfig`]: config::PipelineConfig
105//! [`ModelRouter`]: routing::router::ModelRouter
106//! [`ArbitrageEngine`]: routing::ArbitrageEngine
107//! [`PoolSizer`]: routing::PoolSizer
108//! [`PromptGuard`]: security::PromptGuard
109
110use std::collections::HashMap;
111use thiserror::Error;
112
113pub mod ab_test;
114pub mod adaptive_pool;
115pub mod adaptive_timeout;
116pub mod cost_estimator;
117pub mod admission_control;
118pub mod circuit_breaker;
119pub mod cache;
120pub mod cascade;
121pub mod compression;
122pub mod config;
123pub mod conversation;
124pub mod conversation_state;
125pub mod eval_harness;
126pub mod coordination;
127pub mod failover;
128pub mod provider_health;
129pub mod rate_limiter;
130pub mod semantic_cache;
131pub mod smart_router;
132#[cfg(feature = "distributed")]
133pub mod distributed;
134pub mod enhanced;
135pub mod multi_pipeline;
136pub mod templates;
137pub mod load_balancer;
138pub mod template;
139pub mod metrics;
140pub mod plugin;
141pub mod request_dedup;
142pub mod routing;
143pub mod scheduler;
144pub mod security;
145pub mod session;
146pub mod session_mgr;
147pub mod job_scheduler;
148pub mod context_mgr;
149pub mod stream_agg;
150pub mod stages;
151pub mod token_budget;
152pub mod worker;
153
154#[cfg(feature = "metrics-server")]
155pub mod metrics_server;
156
157#[cfg(feature = "web-api")]
158pub mod web_api;
159
160#[cfg(feature = "self-tune")]
161pub mod self_tune;
162
163#[cfg(feature = "self-modify")]
164pub mod self_modify;
165
166#[cfg(feature = "intelligence")]
167pub mod intelligence;
168
169#[cfg(feature = "evolution")]
170pub mod evolution;
171
172#[cfg(all(
173 feature = "self-tune",
174 feature = "self-modify",
175 feature = "intelligence"
176))]
177pub mod self_improve;
178
179#[cfg(all(feature = "self-tune", feature = "self-modify"))]
180pub mod self_improve_loop;
181
182#[cfg(feature = "tui")]
183pub mod tui;
184
185pub mod audit;
186pub mod feedback_loop;
187pub mod conversation_graph;
188pub mod pipeline;
189pub mod cost_optimizer;
190pub mod prompt_optimizer;
191pub mod trace_export;
192pub mod trace_ui;
193pub mod webhooks;
194pub mod priority_queue;
195pub mod hot_config;
196pub mod observability;
197pub mod retry_policy;
198pub mod worker_pool;
199pub mod model_selector;
200pub mod intent_classifier;
201pub mod response_validator;
202pub mod prompt_versioning;
203pub mod multi_modal;
204pub mod persona_manager;
205pub mod chain_of_thought;
206pub mod streaming_processor;
207pub mod model_fallback;
208pub mod provider_manager;
209pub mod context_compression;
210pub mod prompt_router;
211pub mod tool_call_parser;
212pub mod session_manager;
213pub mod output_cache;
214pub mod pipeline_builder;
215pub mod prompt_safety;
216pub mod experiment_runner;
217pub mod prompt_template;
218pub mod response_classifier;
219pub mod conversation_analyzer;
220pub mod retry_budget;
221pub mod model_registry;
222pub mod prompt_validator;
223pub mod cache_warmer;
224pub mod token_counter;
225pub mod ab_testing;
226
227// Re-exports
228pub use cache::{CacheConfig, CacheEntry, CacheStats, PromptCache};
229pub use rate_limiter::{RateLimitError, RateLimiterConfig, RateLimiterRegistry, ModelRateLimiter};
230pub use ab_test::{AbTestConfig, AbTestResult, AbTestRunner, SuccessMetric, Variant};
231pub use conversation::{
232 ConversationConfig, ConversationManager, PromptFormat, Role, Turn,
233};
234pub use stages::{
235 spawn_pipeline, spawn_pipeline_with_config, LogSink, OutputSink, PipelineHandles, SinkError,
236};
237pub use templates::{
238 AbExperiment, ExperimentReport, ExperimentVariant, PromptTemplate, TemplateError,
239 TemplateRegistry,
240};
241pub use worker::{
242 stream_worker, AnthropicWorker, EchoWorker, LlamaCppWorker, LoadBalancedWorker, ModelWorker,
243 OpenAiWorker, VllmWorker,
244};
245pub use failover::FailoverChain;
246pub use provider_health::{ProviderHealth, ProviderHealthMonitor};
247pub use smart_router::{ModelPricing, RoutingDecision, RoutingRequirements, SmartRouter};
248pub use load_balancer::{
249 BalancerConfig, EndpointStats, LoadBalancer, LoadBalancerStats, ModelEndpoint,
250};
251pub use template::{TemplateContext, TemplateLibrary, TemplateValue};
252pub use pipeline::{
253 AppendStage, LanguageDetectStage, Pipeline, PipelineBuilder, PipelineError, PipelineResult,
254 PipelineStats, PrependStage, RegexReplaceStage, TrimStage, TruncateStage,
255};
256pub use pipeline::PipelineStage as PromptPipelineStage;
257pub use audit::{AuditEntry, AuditFilter, AuditLog, AuditQueryResponse, AuditStats, AuditStatsResponse};
258
259/// Orchestrator-specific errors.
260///
261/// All variants are non-panicking. Callers should match on the variant to
262/// decide whether to retry, shed, or propagate the error.
263#[derive(Error, Debug)]
264pub enum OrchestratorError {
265 /// A pipeline channel was closed before the request could be delivered.
266 ///
267 /// This typically means a pipeline stage task has exited. The pipeline
268 /// should be restarted. This error is not retryable within the same pipeline instance.
269 #[error("channel closed unexpectedly")]
270 ChannelClosed,
271
272 /// An inference worker returned an error.
273 ///
274 /// The inner string contains the provider error message. May be retryable
275 /// depending on the underlying cause (transient network vs. invalid request).
276 #[error("inference failed: {0}")]
277 Inference(String),
278
279 /// Pipeline or worker configuration is invalid.
280 ///
281 /// Returned during startup validation. Not retryable without a config change.
282 #[error("configuration error: {0}")]
283 ConfigError(String),
284
285 /// Provider returned HTTP 429 — callers should back off for `retry_after`.
286 #[error("rate limited by provider (retry after {retry_after_secs}s)")]
287 RateLimited {
288 /// Seconds to wait before retrying, parsed from `Retry-After` header.
289 retry_after_secs: u64,
290 },
291
292 /// Spending cap reached — no further inference allowed this session.
293 #[error("budget exceeded: spent ${spent:.4} of ${limit:.4} limit")]
294 BudgetExceeded {
295 /// Amount spent so far in USD.
296 spent: f64,
297 /// Configured spending limit in USD.
298 limit: f64,
299 },
300
301 /// Provider rejected the request due to an authentication failure.
302 ///
303 /// Returned when the provider returns HTTP 401 or 403.
304 /// Not retryable without a credential rotation.
305 #[error("authentication failed: {0}")]
306 AuthFailed(String),
307
308 /// The inference call exceeded the configured timeout and was cancelled.
309 ///
310 /// The `timeout_secs` field reflects the timeout that was breached.
311 /// Retryable with a longer timeout or by shedding the request.
312 #[error("inference timed out after {timeout_secs}s")]
313 InferenceTimeout {
314 /// Configured timeout that was exceeded.
315 timeout_secs: u64,
316 },
317
318 /// A catch-all error variant for errors that do not fit the other categories.
319 #[error("{0}")]
320 Other(String),
321}
322
323impl OrchestratorError {
324 /// Return a stable lowercase string identifying the error variant,
325 /// suitable for use as a metric label or structured log field.
326 ///
327 /// # Panics
328 ///
329 /// This function does not panic.
330 pub fn error_kind(&self) -> &'static str {
331 match self {
332 Self::ChannelClosed => "channel_closed",
333 Self::Inference(_) => "inference",
334 Self::ConfigError(_) => "config_error",
335 Self::RateLimited { .. } => "rate_limited",
336 Self::BudgetExceeded { .. } => "budget_exceeded",
337 Self::AuthFailed(_) => "auth_failed",
338 Self::InferenceTimeout { .. } => "inference_timeout",
339 Self::Other(_) => "other",
340 }
341 }
342
343 /// Return `true` if retrying the request may succeed.
344 ///
345 /// Transient errors (`RateLimited`, `InferenceTimeout`, `Inference`) are
346 /// retryable. Permanent errors (`ConfigError`, `AuthFailed`,
347 /// `BudgetExceeded`, `ChannelClosed`) are not.
348 ///
349 /// # Panics
350 ///
351 /// This function does not panic.
352 ///
353 /// # Examples
354 ///
355 /// ```
356 /// use tokio_prompt_orchestrator::OrchestratorError;
357 ///
358 /// assert!(OrchestratorError::RateLimited { retry_after_secs: 5 }.is_retryable());
359 /// assert!(!OrchestratorError::AuthFailed("bad key".into()).is_retryable());
360 /// ```
361 pub fn is_retryable(&self) -> bool {
362 matches!(
363 self,
364 Self::RateLimited { .. } | Self::InferenceTimeout { .. } | Self::Inference(_)
365 )
366 }
367}
368
369/// Unique session identifier for request tracking and affinity
370#[derive(Debug, Clone, PartialEq, Eq, Hash, serde::Serialize, serde::Deserialize)]
371pub struct SessionId(pub String);
372
373impl SessionId {
374 /// Create a new `SessionId` from any string-like value.
375 ///
376 /// # Panics
377 ///
378 /// This function does not panic.
379 pub fn new(id: impl Into<String>) -> Self {
380 Self(id.into())
381 }
382
383 /// Borrow the inner string slice.
384 ///
385 /// # Panics
386 ///
387 /// This function does not panic.
388 pub fn as_str(&self) -> &str {
389 &self.0
390 }
391}
392
393impl std::fmt::Display for SessionId {
394 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
395 f.write_str(&self.0)
396 }
397}
398
399/// Initial prompt request from client
400#[derive(Debug, Clone)]
401pub struct PromptRequest {
402 /// Session this request belongs to. Used for affinity sharding and
403 /// deduplication key generation.
404 pub session: SessionId,
405 /// Unique ID for distributed trace correlation across all pipeline stages.
406 pub request_id: String,
407 /// The raw prompt text to send to the inference backend.
408 pub input: String,
409 /// Arbitrary key-value metadata forwarded unchanged through the pipeline.
410 pub meta: HashMap<String, String>,
411 /// Optional absolute deadline for this request. When `Some`, the inference
412 /// stage drops the request and increments `requests_expired_total` if the
413 /// deadline has already passed at dequeue time.
414 pub deadline: Option<std::time::Instant>,
415}
416
417impl PromptRequest {
418 /// Builder-style helper that sets an absolute deadline `duration` from now.
419 ///
420 /// This is an infallible convenience wrapper around `try_with_deadline`.
421 /// It accepts any positive `Duration`; for validated input (e.g. user-supplied
422 /// values) prefer `try_with_deadline` which rejects zero or excessively large
423 /// timeouts at the call site.
424 ///
425 /// # Example
426 ///
427 /// ```
428 /// use std::collections::HashMap;
429 /// use std::time::Duration;
430 /// use tokio_prompt_orchestrator::{PromptRequest, SessionId};
431 ///
432 /// let req = PromptRequest {
433 /// session: SessionId::new("s1"),
434 /// request_id: "r1".to_string(),
435 /// input: "hello".to_string(),
436 /// meta: HashMap::new(),
437 /// deadline: None,
438 /// }
439 /// .with_deadline(Duration::from_secs(5));
440 ///
441 /// assert!(req.deadline.is_some());
442 /// ```
443 ///
444 /// # Panics
445 ///
446 /// This function does not panic.
447 #[must_use]
448 pub fn with_deadline(mut self, duration: std::time::Duration) -> Self {
449 self.deadline = Some(std::time::Instant::now() + duration);
450 self
451 }
452
453 /// Validated builder that sets an absolute deadline `timeout_seconds` from now.
454 ///
455 /// Accepts a timeout expressed as a whole number of seconds and validates it
456 /// before storing the deadline.
457 ///
458 /// # Errors
459 ///
460 /// Returns `Err(OrchestratorError::ConfigError(_))` when:
461 /// - `timeout_seconds` is `0` — a zero-second deadline expires immediately.
462 /// - `timeout_seconds` exceeds `3600` — unreasonably large timeouts are
463 /// rejected to prevent accidental resource leaks.
464 ///
465 /// # Example
466 ///
467 /// ```
468 /// use std::collections::HashMap;
469 /// use tokio_prompt_orchestrator::{PromptRequest, SessionId};
470 ///
471 /// let req = PromptRequest {
472 /// session: SessionId::new("s1"),
473 /// request_id: "r1".to_string(),
474 /// input: "hello".to_string(),
475 /// meta: HashMap::new(),
476 /// deadline: None,
477 /// }
478 /// .try_with_deadline(30)
479 /// .expect("30s is a valid timeout");
480 ///
481 /// assert!(req.deadline.is_some());
482 /// ```
483 ///
484 /// # Panics
485 ///
486 /// This function does not panic.
487 pub fn try_with_deadline(
488 mut self,
489 timeout_seconds: u64,
490 ) -> Result<Self, OrchestratorError> {
491 if timeout_seconds == 0 {
492 return Err(OrchestratorError::ConfigError(
493 "timeout_seconds must be > 0".into(),
494 ));
495 }
496 if timeout_seconds > 3600 {
497 return Err(OrchestratorError::ConfigError(
498 "timeout_seconds must be \u{2264} 3600".into(),
499 ));
500 }
501 self.deadline = Some(
502 std::time::Instant::now()
503 + std::time::Duration::from_secs(timeout_seconds),
504 );
505 Ok(self)
506 }
507}
508
509/// Output from the RAG (retrieval-augmented generation) stage.
510///
511/// Carries the retrieved context string alongside the original request so the
512/// assembly stage can compose the final prompt without re-reading the original.
513#[derive(Debug, Clone)]
514pub struct RagOutput {
515 /// Session this output belongs to.
516 pub session: SessionId,
517 /// Retrieved context text (documents, embeddings, etc.) for the prompt.
518 pub context: String,
519 /// The original request, forwarded intact for use in the assembly stage.
520 pub original: PromptRequest,
521 /// Deadline propagated from the originating [`PromptRequest`], if any.
522 pub deadline: Option<std::time::Instant>,
523}
524
525/// Output from the prompt assembly stage.
526///
527/// The assembled `prompt` string is the final, context-injected input that will
528/// be sent to the inference worker.
529#[derive(Debug, Clone)]
530pub struct AssembleOutput {
531 /// Session this output belongs to.
532 pub session: SessionId,
533 /// Unique request identifier for distributed trace correlation.
534 pub request_id: String,
535 /// The fully assembled prompt string ready for inference.
536 pub prompt: String,
537 /// Deadline propagated from the originating [`PromptRequest`], if any.
538 pub deadline: Option<std::time::Instant>,
539}
540
541/// Output from the inference stage.
542///
543/// Contains the raw token list as returned by the model worker before
544/// post-processing joins them into a coherent response string.
545#[derive(Debug, Clone)]
546pub struct InferenceOutput {
547 /// Session this output belongs to.
548 pub session: SessionId,
549 /// Unique request identifier for distributed trace correlation.
550 pub request_id: String,
551 /// Raw token list from the model worker.
552 pub tokens: Vec<String>,
553}
554
555/// Output from the post-processing stage.
556///
557/// Tokens have been joined, filtered, and formatted into the final
558/// response string delivered to the stream stage.
559#[derive(Debug, Clone)]
560pub struct PostOutput {
561 /// Session this output belongs to.
562 pub session: SessionId,
563 /// Unique request identifier for distributed trace correlation.
564 pub request_id: String,
565 /// Final response text after post-processing.
566 pub text: String,
567}
568
569/// FNV-1a hash — deterministic across process restarts unlike `DefaultHasher`.
570fn fnv1a_hash(s: &str) -> u64 {
571 const PRIME: u64 = 1_099_511_628_211;
572 const BASIS: u64 = 14_695_981_039_346_656_037;
573 s.bytes()
574 .fold(BASIS, |acc, b| acc.wrapping_mul(PRIME) ^ b as u64)
575}
576
577/// Session affinity sharding helper.
578///
579/// Uses FNV-1a for stable hashing across process restarts so that the same
580/// session always routes to the same shard after a restart.
581///
582/// # Panics
583///
584/// This function does not panic.
585pub fn shard_session(session: &SessionId, shards: usize) -> usize {
586 if shards == 0 {
587 return 0;
588 }
589 (fnv1a_hash(&session.0) as usize) % shards
590}
591
592/// A request that was dropped (shed) by the pipeline due to backpressure or
593/// failure. Stored in the [`DeadLetterQueue`] for inspection and replay.
594#[derive(Debug, Clone)]
595pub struct DroppedRequest {
596 /// The original request ID for trace correlation.
597 pub request_id: String,
598 /// The session this request belonged to.
599 pub session_id: String,
600 /// Human-readable reason the request was dropped.
601 pub reason: String,
602 /// Wall-clock time at which the request was dropped.
603 pub dropped_at: std::time::SystemTime,
604}
605
606/// In-memory dead-letter queue for shed pipeline requests.
607///
608/// Stores up to `capacity` most-recent dropped requests in a ring buffer.
609/// When full, the oldest entry is evicted to make room for the newest.
610///
611/// Thread-safe via an internal `Mutex`. Clone is cheap — all clones share
612/// the same underlying ring.
613#[derive(Clone)]
614pub struct DeadLetterQueue {
615 inner: std::sync::Arc<std::sync::Mutex<std::collections::VecDeque<DroppedRequest>>>,
616 capacity: usize,
617}
618
619impl DeadLetterQueue {
620 /// Create a new `DeadLetterQueue` with the given ring-buffer capacity.
621 ///
622 /// # Panics
623 ///
624 /// This function does not panic.
625 pub fn new(capacity: usize) -> Self {
626 Self {
627 inner: std::sync::Arc::new(std::sync::Mutex::new(
628 std::collections::VecDeque::with_capacity(capacity.min(1024)),
629 )),
630 capacity,
631 }
632 }
633
634 /// Push a dropped request into the queue. Evicts the oldest entry if full.
635 ///
636 /// # Panics
637 ///
638 /// This function does not panic. If the internal mutex is poisoned it is
639 /// recovered automatically and a warning is logged.
640 pub fn push(&self, req: DroppedRequest) {
641 // NOTE: std::sync::Mutex is safe here because the critical section is
642 // extremely short (a deque push plus an optional pop_front) and there
643 // is no `.await` inside the guard. Using a sync lock avoids the
644 // overhead of a tokio::sync::Mutex while keeping the operation
645 // compatible with both sync and async callers.
646 let mut guard = self.inner.lock().unwrap_or_else(|p| {
647 tracing::warn!("DeadLetterQueue: recovering from poisoned mutex");
648 crate::metrics::inc_dlq_lock_poisoned();
649 p.into_inner()
650 });
651 if guard.len() >= self.capacity {
652 guard.pop_front();
653 }
654 guard.push_back(req);
655 }
656
657 /// Drain all queued entries and return them, clearing the queue.
658 ///
659 /// # Panics
660 ///
661 /// This function does not panic. If the internal mutex is poisoned it is
662 /// recovered automatically and a warning is logged.
663 pub fn drain(&self) -> Vec<DroppedRequest> {
664 // NOTE: std::sync::Mutex is safe here because the critical section is
665 // a single drain-and-collect with no `.await` inside the guard. The
666 // lock is always released before any async work can be scheduled.
667 let mut guard = self.inner.lock().unwrap_or_else(|p| {
668 tracing::warn!("DeadLetterQueue: recovering from poisoned mutex");
669 crate::metrics::inc_dlq_lock_poisoned();
670 p.into_inner()
671 });
672 guard.drain(..).collect()
673 }
674
675 /// Return the number of entries currently in the queue.
676 ///
677 /// # Panics
678 ///
679 /// This function does not panic. If the internal mutex is poisoned it is
680 /// recovered automatically and a warning is logged.
681 pub fn len(&self) -> usize {
682 self.inner
683 .lock()
684 .unwrap_or_else(|p| {
685 tracing::warn!("DeadLetterQueue: recovering from poisoned mutex");
686 crate::metrics::inc_dlq_lock_poisoned();
687 p.into_inner()
688 })
689 .len()
690 }
691
692 /// Return `true` if the queue contains no entries.
693 ///
694 /// # Panics
695 ///
696 /// This function does not panic. If the internal mutex is poisoned it is
697 /// recovered automatically and a warning is logged.
698 pub fn is_empty(&self) -> bool {
699 self.len() == 0
700 }
701
702 /// Return the maximum number of entries this queue can hold before evicting.
703 ///
704 /// # Panics
705 ///
706 /// This function does not panic.
707 pub fn capacity(&self) -> usize {
708 self.capacity
709 }
710
711 /// Return a snapshot of all queued entries without removing them.
712 ///
713 /// Unlike [`drain`](Self::drain), this does not clear the queue. The
714 /// snapshot is a clone taken under the lock; the queue continues to
715 /// operate normally while the returned `Vec` is used.
716 ///
717 /// # Panics
718 ///
719 /// This function does not panic. If the internal mutex is poisoned it is
720 /// recovered automatically and a warning is logged.
721 pub fn peek(&self) -> Vec<DroppedRequest> {
722 self.inner
723 .lock()
724 .unwrap_or_else(|p| {
725 tracing::warn!("DeadLetterQueue: recovering from poisoned mutex");
726 crate::metrics::inc_dlq_lock_poisoned();
727 p.into_inner()
728 })
729 .iter()
730 .cloned()
731 .collect()
732 }
733}
734
735/// Type alias for the optional OpenTelemetry tracing layer used in main.rs.
736///
737/// Exported so binary crates can declare `Option<OtelLayer>` without spelling
738/// out the full generic type.
739pub type OtelLayer = tracing_opentelemetry::OpenTelemetryLayer<
740 tracing_subscriber::Registry,
741 opentelemetry_sdk::trace::Tracer,
742>;
743
744/// Initialise tracing with env-filter support. Call once at binary startup.
745///
746/// # Panics
747///
748/// This function does not panic.
749///
750/// ## Log format
751///
752/// If `RUST_LOG_FORMAT=json` is set the subscriber emits newline-delimited
753/// JSON suitable for log aggregation pipelines. Otherwise the human-readable
754/// `fmt` pretty format is used for local development.
755///
756/// ## OpenTelemetry OTLP export
757///
758/// If the environment variable `OTEL_EXPORTER_OTLP_ENDPOINT` (or the legacy
759/// `JAEGER_ENDPOINT`) is set to a valid OTLP collector URL (e.g.
760/// `http://localhost:4318`), spans are exported via OTLP HTTP to that endpoint
761/// using a batch exporter on the Tokio runtime.
762///
763/// If neither variable is set (the common case in local development), the OTel
764/// layer is **silently omitted** — no error is printed and the binary starts
765/// normally. All other tracing output (stdout / log files) is unaffected.
766///
767/// A startup `info!` log is emitted either way so operators can confirm the
768/// observability configuration at a glance:
769///
770/// - `"OpenTelemetry OTLP export enabled, sending to <endpoint>"`
771/// - `"OpenTelemetry OTLP disabled (set OTEL_EXPORTER_OTLP_ENDPOINT to enable)"`
772///
773/// ## Calling requirement
774///
775/// This function **must** be called after the Tokio runtime has started because
776/// the OTLP batch exporter uses `rt-tokio` internally.
777pub fn init_tracing() {
778 use tracing_subscriber::{fmt, layer::SubscriberExt, EnvFilter, Layer, Registry};
779
780 let env_filter = EnvFilter::try_from_default_env().unwrap_or_else(|_| EnvFilter::new("info"));
781
782 let use_json = std::env::var("RUST_LOG_FORMAT")
783 .map(|v| v.to_lowercase() == "json")
784 .unwrap_or(false);
785
786 // Build the OTel layer if an endpoint is configured, boxing it so the
787 // concrete type does not propagate into the subscriber stack.
788 let otel_layer: Option<Box<dyn Layer<Registry> + Send + Sync>> =
789 try_build_otel_layer().map(|l| l.boxed());
790
791 if use_json {
792 let subscriber = Registry::default()
793 .with(otel_layer)
794 .with(env_filter)
795 .with(fmt::layer().json().with_target(false));
796 let _ = tracing::subscriber::set_global_default(subscriber);
797 } else {
798 let subscriber = Registry::default()
799 .with(otel_layer)
800 .with(env_filter)
801 .with(fmt::layer().with_target(false));
802 let _ = tracing::subscriber::set_global_default(subscriber);
803 }
804}
805
806/// Attempt to build an OpenTelemetry tracing layer, returning `None` on error.
807///
808/// Reads `JAEGER_ENDPOINT` (e.g. `http://localhost:4318`) or
809/// `OTEL_EXPORTER_OTLP_ENDPOINT`. When neither is set, or when the exporter
810/// fails to build, this function returns `None` so startup is never blocked
811/// by observability infrastructure.
812///
813/// **NOTE**: This function must be called **after** the Tokio runtime has been
814/// started because the batch exporter uses `rt-tokio`.
815///
816/// # Panics
817///
818/// This function never panics.
819pub fn try_build_otel_layer() -> Option<
820 tracing_opentelemetry::OpenTelemetryLayer<
821 tracing_subscriber::Registry,
822 opentelemetry_sdk::trace::Tracer,
823 >,
824> {
825 use opentelemetry::global;
826 use opentelemetry_otlp::WithExportConfig;
827
828 let endpoint = std::env::var("OTEL_EXPORTER_OTLP_ENDPOINT")
829 .or_else(|_| std::env::var("JAEGER_ENDPOINT"))
830 .ok();
831
832 let endpoint = match endpoint {
833 Some(ep) => {
834 tracing::info!(
835 endpoint = ep.as_str(),
836 "OpenTelemetry OTLP export enabled, sending to {ep}"
837 );
838 ep
839 }
840 None => {
841 tracing::info!(
842 "OpenTelemetry OTLP disabled (set OTEL_EXPORTER_OTLP_ENDPOINT to enable)"
843 );
844 return None;
845 }
846 };
847
848 let exporter = match opentelemetry_otlp::SpanExporter::builder()
849 .with_http()
850 .with_endpoint(endpoint)
851 .build()
852 {
853 Ok(e) => e,
854 Err(e) => {
855 tracing::warn!("OTel exporter build failed: {e}");
856 return None;
857 }
858 };
859
860 let provider = opentelemetry_sdk::trace::TracerProvider::builder()
861 .with_batch_exporter(exporter, opentelemetry_sdk::runtime::Tokio)
862 .with_resource(opentelemetry_sdk::Resource::new(vec![
863 opentelemetry::KeyValue::new("service.name", "tokio-prompt-orchestrator"),
864 ]))
865 .build();
866
867 use opentelemetry::trace::TracerProvider as _;
868 let tracer = provider.tracer("tokio-prompt-orchestrator");
869 global::set_tracer_provider(provider);
870 Some(tracing_opentelemetry::layer().with_tracer(tracer))
871}
872
873/// Outcome of a [`send_with_shed`] call.
874///
875/// Distinguishes between successful delivery and a graceful shed so callers
876/// can log/metric them differently.
877#[derive(Debug, Clone, Copy, PartialEq, Eq)]
878pub enum SendOutcome {
879 /// The item was successfully placed in the channel.
880 Queued,
881 /// The channel was full; the item was dropped to shed load.
882 Shed,
883}
884
885/// Pipeline stage identifier for metrics and logging.
886#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
887pub enum PipelineStage {
888 /// Retrieval-augmented generation — fetches context before assembly.
889 Rag,
890 /// Prompt assembly — combines context and user input into a full prompt.
891 Assemble,
892 /// Inference — sends the assembled prompt to the model backend.
893 Inference,
894 /// Post-processing — formats and filters the raw model response.
895 Post,
896 /// Streaming output — delivers the processed response downstream.
897 Stream,
898}
899
900impl PipelineStage {
901 /// Return the canonical lowercase ASCII name used in metrics labels.
902 ///
903 /// # Panics
904 ///
905 /// This function does not panic.
906 pub fn as_str(&self) -> &'static str {
907 match self {
908 Self::Rag => "rag",
909 Self::Assemble => "assemble",
910 Self::Inference => "inference",
911 Self::Post => "post",
912 Self::Stream => "stream",
913 }
914 }
915
916 /// Return a slice of all pipeline stage variants in pipeline order.
917 ///
918 /// Useful for iterating over all stages when initialising per-stage metrics
919 /// or building dashboards.
920 ///
921 /// # Examples
922 ///
923 /// ```
924 /// use tokio_prompt_orchestrator::PipelineStage;
925 ///
926 /// let labels: Vec<&str> = PipelineStage::all().iter().map(|s| s.as_str()).collect();
927 /// assert_eq!(labels, ["rag", "assemble", "inference", "post", "stream"]);
928 /// ```
929 ///
930 /// # Panics
931 ///
932 /// This function does not panic.
933 pub fn all() -> &'static [PipelineStage] {
934 &[
935 PipelineStage::Rag,
936 PipelineStage::Assemble,
937 PipelineStage::Inference,
938 PipelineStage::Post,
939 PipelineStage::Stream,
940 ]
941 }
942}
943
944impl std::fmt::Display for PipelineStage {
945 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
946 f.write_str(self.as_str())
947 }
948}
949
950/// Send with graceful shedding on backpressure.
951///
952/// # Returns
953/// - `Ok(SendOutcome::Queued)` — item was accepted by the channel.
954/// - `Ok(SendOutcome::Shed)` — channel was full; item was dropped gracefully.
955/// - `Err(OrchestratorError::ChannelClosed)` — the receiver has been dropped.
956///
957/// # Panics
958///
959/// This function never panics.
960pub async fn send_with_shed<T>(
961 tx: &tokio::sync::mpsc::Sender<T>,
962 item: T,
963 stage: PipelineStage,
964) -> Result<SendOutcome, OrchestratorError> {
965 match tx.try_send(item) {
966 Ok(_) => Ok(SendOutcome::Queued),
967 Err(tokio::sync::mpsc::error::TrySendError::Full(_)) => {
968 tracing::warn!(stage = %stage, "queue full, shedding request");
969 Ok(SendOutcome::Shed)
970 }
971 Err(tokio::sync::mpsc::error::TrySendError::Closed(_)) => {
972 Err(OrchestratorError::ChannelClosed)
973 }
974 }
975}
976
977#[cfg(test)]
978mod tests {
979 use super::*;
980
981 #[tokio::test]
982 async fn test_send_with_shed_returns_shed_outcome_when_channel_full() {
983 // Channel of capacity 1 — fill it, then send again to trigger shed
984 let (tx, _rx) = tokio::sync::mpsc::channel::<u32>(1);
985 // Fill the channel
986 let first = send_with_shed(&tx, 1u32, PipelineStage::Rag).await.unwrap();
987 assert_eq!(first, SendOutcome::Queued, "first send should be Queued");
988 // Channel is now full — next send must be Shed
989 let second = send_with_shed(&tx, 2u32, PipelineStage::Rag).await.unwrap();
990 assert_eq!(
991 second,
992 SendOutcome::Shed,
993 "second send on full channel should be Shed"
994 );
995 }
996
997 #[tokio::test]
998 async fn test_send_with_shed_returns_queued_when_space_available() {
999 let (tx, _rx) = tokio::sync::mpsc::channel::<u32>(10);
1000 let outcome = send_with_shed(&tx, 42u32, PipelineStage::Rag)
1001 .await
1002 .unwrap();
1003 assert_eq!(outcome, SendOutcome::Queued, "send into channel with space must return Queued");
1004 }
1005
1006 #[tokio::test]
1007 async fn test_send_with_shed_returns_error_when_channel_closed() {
1008 let (tx, rx) = tokio::sync::mpsc::channel::<u32>(10);
1009 drop(rx);
1010 let result = send_with_shed(&tx, 1u32, PipelineStage::Rag).await;
1011 assert!(matches!(result, Err(OrchestratorError::ChannelClosed)), "send after receiver drop must return ChannelClosed error");
1012 }
1013
1014 #[test]
1015 fn test_shard_session_deterministic() {
1016 let s = SessionId::new("test-session-123");
1017 assert_eq!(shard_session(&s, 4), shard_session(&s, 4), "shard_session must be deterministic for the same session and shard count");
1018 }
1019
1020 #[test]
1021 fn test_shard_session_distribution() {
1022 let sessions: Vec<_> = (0..100).map(|i| SessionId::new(format!("s-{i}"))).collect();
1023 let counts: Vec<_> = (0..4usize)
1024 .map(|sh| {
1025 sessions
1026 .iter()
1027 .filter(|s| shard_session(s, 4) == sh)
1028 .count()
1029 })
1030 .collect();
1031 assert!(counts.iter().all(|&c| c > 0), "each of the 4 shards must receive at least one session from the 100-session sample");
1032 }
1033
1034 /// Verify that tracing events can be captured and that when RUST_LOG_FORMAT=json
1035 /// is set, the output is valid newline-delimited JSON with expected fields.
1036 #[test]
1037 fn test_json_log_output_is_valid_json() {
1038 use std::sync::{Arc, Mutex};
1039 use tracing_subscriber::{fmt, layer::SubscriberExt, EnvFilter};
1040
1041 // Shared buffer to capture log output.
1042 let buf: Arc<Mutex<Vec<u8>>> = Arc::new(Mutex::new(Vec::new()));
1043 let buf_clone = buf.clone();
1044
1045 // Build an isolated JSON subscriber for this test.
1046 let writer = tracing_subscriber::fmt::writer::BoxMakeWriter::new(move || {
1047 struct BufWriter(Arc<Mutex<Vec<u8>>>);
1048 impl std::io::Write for BufWriter {
1049 fn write(&mut self, b: &[u8]) -> std::io::Result<usize> {
1050 self.0
1051 .lock()
1052 .unwrap_or_else(|p| p.into_inner())
1053 .extend_from_slice(b);
1054 Ok(b.len())
1055 }
1056 fn flush(&mut self) -> std::io::Result<()> {
1057 Ok(())
1058 }
1059 }
1060 BufWriter(buf_clone.clone())
1061 });
1062
1063 let subscriber = tracing_subscriber::Registry::default()
1064 .with(EnvFilter::new("info"))
1065 .with(fmt::layer().json().with_writer(writer).with_target(false));
1066
1067 // Use a local dispatcher so this test doesn't interfere with globals.
1068 let _guard = tracing::subscriber::with_default(subscriber, || {
1069 tracing::info!(
1070 stage = "rag",
1071 session_id = "s1",
1072 request_id = "r1",
1073 "test event"
1074 );
1075 });
1076
1077 let captured = buf.lock().unwrap_or_else(|p| p.into_inner()).clone();
1078 assert!(!captured.is_empty(), "captured log must be non-empty");
1079
1080 let text = std::str::from_utf8(&captured).expect("log output must be valid UTF-8");
1081 // Each line must be valid JSON.
1082 for line in text.lines().filter(|l| !l.is_empty()) {
1083 let v: serde_json::Value =
1084 serde_json::from_str(line).expect("each log line must be valid JSON");
1085 // Verify expected fields are present.
1086 assert!(v.get("fields").is_some(), "JSON log must have 'fields' key");
1087 }
1088 }
1089
1090 /// Verify that trace IDs remain consistent across a pipeline processing a
1091 /// single request (OTel context propagation). Without a live collector the
1092 /// test only checks that the tracing infrastructure works without panicking;
1093 /// the trace_id field is non-zero within a span.
1094 #[test]
1095 fn test_trace_id_is_non_zero_within_span() {
1096 use opentelemetry::trace::{SpanContext, TraceContextExt};
1097 use tracing_opentelemetry::OpenTelemetrySpanExt;
1098 use tracing_subscriber::layer::SubscriberExt;
1099
1100 // Set up a minimal OTel provider with default config (no exporter).
1101 let provider = opentelemetry_sdk::trace::TracerProvider::builder().build();
1102 use opentelemetry::trace::TracerProvider as _;
1103 let tracer = provider.tracer("test");
1104
1105 let subscriber = tracing_subscriber::Registry::default()
1106 .with(tracing_opentelemetry::layer().with_tracer(tracer));
1107
1108 tracing::subscriber::with_default(subscriber, || {
1109 let span = tracing::info_span!("test.root");
1110 let _guard = span.enter();
1111 let ctx = tracing::Span::current().context();
1112 let span_ref = ctx.span();
1113 let span_ctx: &SpanContext = span_ref.span_context();
1114 // trace_id must be non-zero when within a valid span.
1115 assert!(
1116 span_ctx.is_valid(),
1117 "span context must be valid inside an instrumented span"
1118 );
1119 assert_ne!(
1120 span_ctx.trace_id(),
1121 opentelemetry::trace::TraceId::INVALID,
1122 "trace_id must be non-zero"
1123 );
1124 });
1125 }
1126
1127 #[test]
1128 fn test_prompt_request_with_deadline_sets_future_instant() {
1129 use std::time::Duration;
1130
1131 let req = PromptRequest {
1132 session: SessionId::new("s1"),
1133 request_id: "r1".to_string(),
1134 input: "hello".to_string(),
1135 meta: HashMap::new(),
1136 deadline: None,
1137 }
1138 .with_deadline(Duration::from_secs(10));
1139
1140 let deadline = req
1141 .deadline
1142 .expect("deadline must be Some after with_deadline");
1143 assert!(
1144 deadline > std::time::Instant::now(),
1145 "deadline must be in the future"
1146 );
1147 }
1148
1149 #[test]
1150 fn test_prompt_request_default_deadline_is_none() {
1151 let req = PromptRequest {
1152 session: SessionId::new("s2"),
1153 request_id: "r2".to_string(),
1154 input: "world".to_string(),
1155 meta: HashMap::new(),
1156 deadline: None,
1157 };
1158 assert!(req.deadline.is_none(), "default deadline must be None");
1159 }
1160}