Skip to main content

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//! ![demo](https://raw.githubusercontent.com/Mattbusel/tokio-prompt-orchestrator/main/assets/demo.gif)
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}