Skip to main content

tokio_prompt_orchestrator/config/
mod.rs

1//! # Stage: Declarative Pipeline Configuration
2//!
3//! ## Responsibility
4//! Parse, validate, and hot-reload TOML pipeline configuration files.
5//! Users define an entire pipeline topology declaratively and run it with:
6//! ```text
7//! cargo run -- --config pipeline.toml
8//! ```
9//!
10//! ## Guarantees
11//! - Deterministic: same TOML input always produces the same `PipelineConfig`
12//! - Validated: all semantic constraints are checked before a config is accepted
13//! - Type-safe: invalid field combinations are caught at parse time via serde
14//! - Hot-reloadable: file changes are detected and validated before applying
15//! - Schema-exportable: JSON Schema output enables IDE autocomplete
16//!
17//! ## NOT Responsible For
18//! - Building the runtime pipeline from config (that belongs to `stages`)
19//! - Managing worker connections (that belongs to `worker`)
20//! - Metrics collection (that belongs to `metrics`)
21
22pub mod loader;
23pub mod validation;
24pub mod watcher;
25
26#[cfg(feature = "schema")]
27use schemars::JsonSchema;
28use serde::{Deserialize, Serialize};
29
30// ── Default value functions ──────────────────────────────────────────────
31
32/// Default RAG stage timeout: 5000 ms.
33///
34/// **Why 5 s?** Covers P99 of most retrieval backends (vector DBs, BM25 indexes)
35/// under normal load while still bounding worst-case pipeline latency.  Increase
36/// for slow cross-region retrieval or when using large embedding models for
37/// re-ranking.
38fn default_timeout_ms() -> u64 {
39    5000
40}
41
42/// Default maximum context tokens for the RAG stage: 2048 tokens.
43///
44/// **Why 2048?** Fits comfortably in the context window of all major models
45/// (≥4K context) while leaving ample room for the instruction, user prompt,
46/// and model response.  Increase to 4096–8192 for RAG-heavy workloads where
47/// retrieved passages tend to be long.
48fn default_max_context_tokens() -> usize {
49    2048
50}
51
52/// Default retry base delay: 100 ms.
53///
54/// **Why 100 ms?** With exponential back-off (2×) this reaches ~3 s by
55/// attempt 5, covering most transient LLM API failures (rate-limit windows,
56/// brief network blips) without introducing multi-second latency on the first
57/// retry.  Reduce to 50 ms for latency-sensitive workflows; increase for
58/// providers with aggressive rate limits.
59fn default_retry_base_ms() -> u64 {
60    100
61}
62
63/// Default retry maximum delay: 5000 ms.
64///
65/// **Why 5 s?** Caps the exponential back-off ceiling so no single retry ever
66/// blocks a pipeline stage for more than 5 s.  This balances giving slow
67/// providers time to recover with keeping end-to-end P99 latency predictable.
68fn default_retry_max_ms() -> u64 {
69    5000
70}
71
72/// Default deduplication window: 300 seconds (5 minutes).
73///
74/// **Why 5 minutes?** Covers typical LLM retry storms — clients that retry on
75/// error within a few minutes will hit the in-flight dedup cache rather than
76/// issuing a duplicate API call.  Longer windows reduce redundant API costs
77/// but proportionally increase memory use; at 10K entries × 5 min the overhead
78/// is negligible.
79fn default_dedup_window_s() -> u64 {
80    300
81}
82
83/// Default deduplication max entries: 10 000.
84///
85/// **Why 10 000?** At ~256 bytes per entry (key hash + response pointer) this
86/// costs ~2.5 MB of heap, which is acceptable even in memory-constrained
87/// deployments.  For high-cardinality workloads (many unique sessions) increase
88/// to 50K–100K; for embedded deployments reduce to 1K–5K.
89fn default_dedup_max_entries() -> usize {
90    10_000
91}
92
93/// Default channel capacity for pipeline stages: 512 entries.
94///
95/// **Why 512?** Sized for ~50 ms of burst traffic at 10 000 req/s while
96/// keeping per-stage memory under ~4 MB (each `PromptRequest` ≈ 8 KB).
97/// This gives downstream stages enough headroom to absorb a slow GC pause or
98/// brief processing spike without shedding load.  Reduce to 128–256 for
99/// memory-constrained deployments; increase to 1024–2048 for very high
100/// throughput or bursty workloads.
101fn default_channel_capacity() -> usize {
102    512
103}
104
105/// Default enabled state: `true`.
106///
107/// **Why true?** Most stages should be on by default so that a minimal config
108/// file (with only required fields) produces a fully-functional pipeline.
109/// Stages that should be off by default (rate limiting) have their own
110/// `default_*_enabled` functions that return `false`.
111fn default_true() -> bool {
112    true
113}
114
115// ── Channel sizes ────────────────────────────────────────────────────────
116
117/// Per-channel capacity overrides for the five inter-stage pipeline channels.
118///
119/// All fields are optional; when absent the stage default is used
120/// (512, 512, 1024, 512, 256 respectively).
121///
122/// # Panics
123///
124/// This type never panics.
125#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
126#[cfg_attr(feature = "schema", derive(JsonSchema))]
127pub struct ChannelSizes {
128    /// Capacity of the RAG → Assemble channel (default: 512).
129    pub rag_to_assemble: Option<usize>,
130    /// Capacity of the Assemble → Inference channel (default: 512).
131    pub assemble_to_inference: Option<usize>,
132    /// Capacity of the Inference → Post channel (default: 1024).
133    pub inference_to_post: Option<usize>,
134    /// Capacity of the Post → Stream channel (default: 512).
135    pub post_to_stream: Option<usize>,
136    /// Capacity of the Stream output channel (default: 256).
137    pub stream_output: Option<usize>,
138}
139
140// ── Top-level config ─────────────────────────────────────────────────────
141
142/// Root configuration for a pipeline instance.
143///
144/// Deserialized from a TOML file and validated before use.
145/// Every field has either a required value or a documented default.
146///
147/// # Example
148///
149/// ```toml
150/// [pipeline]
151/// name = "production"
152/// version = "1.0"
153///
154/// [stages.inference]
155/// worker = "open_ai"
156/// model = "gpt-4"
157/// ```
158///
159/// # Panics
160///
161/// This type never panics during construction or access.
162#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
163#[cfg_attr(feature = "schema", derive(JsonSchema))]
164#[serde(deny_unknown_fields)]
165pub struct PipelineConfig {
166    /// Pipeline identity and version metadata.
167    pub pipeline: PipelineSection,
168    /// Per-stage configuration for the five pipeline stages.
169    pub stages: StagesConfig,
170    /// Resilience settings: retries, circuit breaker, backpressure.
171    pub resilience: ResilienceConfig,
172    /// Rate limiting settings to control request throughput.
173    #[serde(default)]
174    pub rate_limits: RateLimitConfig,
175    /// Deduplication settings for request coalescing.
176    pub deduplication: DeduplicationConfig,
177    /// Observability: logging, metrics, tracing.
178    pub observability: ObservabilityConfig,
179    /// Optional per-channel capacity overrides.  When absent, each stage uses
180    /// its compiled-in default (512 / 512 / 1024 / 512 / 256).
181    #[serde(default)]
182    pub channel_sizes: Option<ChannelSizes>,
183    /// Optional per-worker configuration overrides.
184    ///
185    /// Allows setting model, API base URL, and connection parameters for each
186    /// worker type directly in the pipeline TOML without modifying worker code.
187    /// Workers read these values at construction time in `spawn_pipeline_with_config`.
188    #[serde(default)]
189    pub workers: WorkersConfig,
190}
191
192// ── Pipeline identity ────────────────────────────────────────────────────
193
194/// Pipeline identity and version metadata.
195///
196/// # Panics
197///
198/// This type never panics.
199#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
200#[cfg_attr(feature = "schema", derive(JsonSchema))]
201#[serde(deny_unknown_fields)]
202pub struct PipelineSection {
203    /// Human-readable pipeline name (e.g., "production", "staging").
204    pub name: String,
205    /// Semantic version of this configuration (e.g., "1.0").
206    pub version: String,
207    /// Optional description for documentation purposes.
208    pub description: Option<String>,
209}
210
211// ── Stage configs ────────────────────────────────────────────────────────
212
213/// Configuration for all five pipeline stages.
214///
215/// # Panics
216///
217/// This type never panics.
218#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
219#[cfg_attr(feature = "schema", derive(JsonSchema))]
220#[serde(deny_unknown_fields)]
221pub struct StagesConfig {
222    /// RAG (retrieval-augmented generation) stage settings.
223    pub rag: RagStageConfig,
224    /// Prompt assembly stage settings.
225    pub assemble: AssembleStageConfig,
226    /// Model inference stage settings.
227    pub inference: InferenceStageConfig,
228    /// Post-processing stage settings.
229    pub post_process: PostProcessStageConfig,
230    /// Output streaming stage settings.
231    pub stream: StreamStageConfig,
232}
233
234/// RAG stage configuration.
235///
236/// Controls retrieval timeout, context limits, and channel sizing.
237///
238/// # Panics
239///
240/// This type never panics.
241#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
242#[cfg_attr(feature = "schema", derive(JsonSchema))]
243#[serde(deny_unknown_fields)]
244pub struct RagStageConfig {
245    /// Whether the RAG stage is enabled. Disabled stages pass-through.
246    #[serde(default = "default_true")]
247    pub enabled: bool,
248    /// Maximum time (ms) to wait for retrieval results.
249    #[serde(default = "default_timeout_ms")]
250    pub timeout_ms: u64,
251    /// Maximum context tokens to prepend from retrieval.
252    #[serde(default = "default_max_context_tokens")]
253    pub max_context_tokens: usize,
254    /// Channel buffer capacity for this stage. `None` uses the pipeline default (512).
255    pub channel_capacity: Option<usize>,
256}
257
258/// Assembly stage configuration.
259///
260/// # Panics
261///
262/// This type never panics.
263#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
264#[cfg_attr(feature = "schema", derive(JsonSchema))]
265#[serde(deny_unknown_fields)]
266pub struct AssembleStageConfig {
267    /// Whether the assembly stage is enabled.
268    #[serde(default = "default_true")]
269    pub enabled: bool,
270    /// Channel buffer capacity for this stage.
271    #[serde(default = "default_channel_capacity")]
272    pub channel_capacity: usize,
273}
274
275/// Default number of concurrent inference workers: 1 (single worker).
276///
277/// **Why 1?** Preserves backward compatibility with existing deployments that
278/// assume a single ordered inference queue.  Increase to the number of
279/// parallel in-flight requests you want to sustain (e.g. 4–8 for high-
280/// throughput deployments with async LLM APIs that support concurrent calls).
281fn default_inference_workers() -> usize {
282    1
283}
284
285/// Default adaptive timeout minimum floor: 500 ms.
286///
287/// **Why 500 ms?** This is a conservative lower bound that prevents the
288/// adaptive algorithm from producing a timeout so short that it triggers
289/// spurious failures during brief periods of low-latency traffic.  Cloud LLM
290/// APIs often have a non-trivial connection establishment cost (~100–300 ms)
291/// that would be mis-classified as a timeout at a sub-500 ms floor.  Operators
292/// running local inference servers with sub-100 ms P99 should lower this to
293/// match their actual minimum latency.
294fn default_adaptive_timeout_min_ms() -> u64 {
295    500
296}
297
298/// Inference stage configuration.
299///
300/// Specifies the worker backend, model, and generation parameters.
301///
302/// # Panics
303///
304/// This type never panics.
305#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
306#[cfg_attr(feature = "schema", derive(JsonSchema))]
307#[serde(deny_unknown_fields)]
308pub struct InferenceStageConfig {
309    /// Which model worker backend to use.
310    pub worker: WorkerKind,
311    /// Model name or identifier (e.g., "gpt-4", "claude-3-opus").
312    pub model: String,
313    /// Maximum tokens to generate. `None` uses the worker default.
314    pub max_tokens: Option<u32>,
315    /// Sampling temperature. `None` uses the worker default.
316    pub temperature: Option<f32>,
317    /// Inference timeout in milliseconds. `None` uses no explicit timeout.
318    pub timeout_ms: Option<u64>,
319    /// Number of concurrent inference worker tasks reading from the same
320    /// input channel.  Default 1 (backward-compatible).  Increase to add
321    /// parallelism for high-throughput deployments.
322    #[serde(default = "default_inference_workers")]
323    pub inference_workers: usize,
324    /// Enable adaptive timeout based on P95 of recent inferences. Overrides timeout_ms if set.
325    #[serde(default)]
326    pub adaptive_timeout_enabled: bool,
327    /// Minimum adaptive timeout floor in milliseconds (only used when adaptive_timeout_enabled = true).
328    #[serde(default = "default_adaptive_timeout_min_ms")]
329    pub adaptive_timeout_min_ms: u64,
330}
331
332/// Supported model worker backends.
333///
334/// # Panics
335///
336/// This type never panics.
337#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
338#[cfg_attr(feature = "schema", derive(JsonSchema))]
339#[serde(deny_unknown_fields)]
340#[serde(rename_all = "snake_case")]
341pub enum WorkerKind {
342    /// OpenAI-compatible API (GPT-4, GPT-3.5, etc.).
343    OpenAi,
344    /// Anthropic Claude API.
345    Anthropic,
346    /// Local llama.cpp server.
347    LlamaCpp,
348    /// vLLM inference server.
349    Vllm,
350    /// Echo worker for testing — returns the prompt as tokens.
351    Echo,
352}
353
354/// Provider identifier enum used in per-worker configuration.
355///
356/// Accepts the following string values in TOML (case-insensitive via custom
357/// `Deserialize`): `"openai"`, `"anthropic"`, `"llama_cpp"` / `"llama-cpp"`,
358/// `"vllm"`, `"echo"`.  Unknown values produce a clear serde error at
359/// parse time.
360///
361/// # Panics
362///
363/// This type never panics.
364#[derive(Debug, Clone, Serialize, PartialEq, Eq)]
365#[cfg_attr(feature = "schema", derive(JsonSchema))]
366pub enum Provider {
367    /// OpenAI-compatible API.
368    OpenAi,
369    /// Anthropic Claude API.
370    Anthropic,
371    /// Local llama.cpp server.
372    LlamaCpp,
373    /// vLLM inference server.
374    Vllm,
375    /// Echo (test) worker.
376    Echo,
377}
378
379impl<'de> serde::Deserialize<'de> for Provider {
380    fn deserialize<D: serde::Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
381        let s = String::deserialize(deserializer)?;
382        match s.to_lowercase().replace('-', "_").as_str() {
383            "openai" => Ok(Provider::OpenAi),
384            "anthropic" => Ok(Provider::Anthropic),
385            "llama_cpp" => Ok(Provider::LlamaCpp),
386            "vllm" => Ok(Provider::Vllm),
387            "echo" => Ok(Provider::Echo),
388            other => Err(serde::de::Error::custom(format!(
389                "unknown provider {other:?}; expected one of: openai, anthropic, llama_cpp, vllm, echo"
390            ))),
391        }
392    }
393}
394
395/// Per-worker connection and model configuration.
396///
397/// Optional overrides for a named worker. All fields are optional; when absent
398/// the worker falls back to its built-in defaults or environment variables.
399///
400/// Workers prefer TOML config over defaults but environment variables still
401/// override TOML values (e.g. `OPENAI_API_KEY` always takes precedence over
402/// any key stored in the config file).
403///
404/// # Example
405///
406/// ```toml
407/// [workers.openai]
408/// model = "gpt-4o"
409/// api_base_url = "https://api.openai.com/v1"
410///
411/// [workers.anthropic]
412/// model = "claude-opus-4-6"
413/// max_tokens = 4096
414///
415/// [workers.llama_cpp]
416/// host = "localhost"
417/// port = 8080
418/// ```
419///
420/// # Panics
421///
422/// This type never panics.
423#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
424#[cfg_attr(feature = "schema", derive(JsonSchema))]
425#[derive(Default)]
426pub struct WorkerConfig {
427    /// Model name to use (e.g. `"gpt-4o"`, `"claude-opus-4-6"`).
428    pub model: Option<String>,
429    /// API base URL (e.g. `"https://api.openai.com/v1"`).
430    /// Overrides the worker's built-in default endpoint.
431    pub api_base_url: Option<String>,
432    /// Maximum tokens to generate.  `None` uses the worker default.
433    pub max_tokens: Option<u32>,
434    /// Sampling temperature.  `None` uses the worker default.
435    pub temperature: Option<f32>,
436    /// Hostname for local servers (llama.cpp / vLLM).
437    pub host: Option<String>,
438    /// Port for local servers (llama.cpp / vLLM).
439    pub port: Option<u16>,
440}
441
442
443/// Per-worker configuration map, keyed by provider name.
444///
445/// Used in the `[workers.*]` TOML sections.
446///
447/// # Panics
448///
449/// This type never panics.
450#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Default)]
451#[cfg_attr(feature = "schema", derive(JsonSchema))]
452pub struct WorkersConfig {
453    /// OpenAI worker configuration overrides.
454    #[serde(default)]
455    pub openai: WorkerConfig,
456    /// Anthropic worker configuration overrides.
457    #[serde(default)]
458    pub anthropic: WorkerConfig,
459    /// llama.cpp worker configuration overrides.
460    #[serde(default)]
461    pub llama_cpp: WorkerConfig,
462    /// vLLM worker configuration overrides.
463    #[serde(default)]
464    pub vllm: WorkerConfig,
465    /// Echo worker configuration overrides (mostly useful for testing).
466    #[serde(default)]
467    pub echo: WorkerConfig,
468}
469
470/// Post-processing stage configuration.
471///
472/// # Panics
473///
474/// This type never panics.
475#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
476#[cfg_attr(feature = "schema", derive(JsonSchema))]
477#[serde(deny_unknown_fields)]
478pub struct PostProcessStageConfig {
479    /// Whether post-processing is enabled.
480    #[serde(default = "default_true")]
481    pub enabled: bool,
482    /// Channel buffer capacity for this stage.
483    #[serde(default = "default_channel_capacity")]
484    pub channel_capacity: usize,
485}
486
487/// Output streaming stage configuration.
488///
489/// # Panics
490///
491/// This type never panics.
492#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
493#[cfg_attr(feature = "schema", derive(JsonSchema))]
494#[serde(deny_unknown_fields)]
495pub struct StreamStageConfig {
496    /// Whether output streaming is enabled.
497    #[serde(default = "default_true")]
498    pub enabled: bool,
499    /// Channel buffer capacity for this stage.
500    #[serde(default = "default_channel_capacity")]
501    pub channel_capacity: usize,
502}
503
504// ── Rate limiting ────────────────────────────────────────────────────────
505
506/// Default requests-per-second limit: 100 req/s.
507///
508/// **Why 100?** A conservative starting point that prevents accidental
509/// saturation of the downstream LLM provider (most have free-tier limits of
510/// ~60 req/min ≈ 1 req/s).  Operators deploying against higher-tier API
511/// accounts should raise this to their actual quota.
512fn default_rps() -> u32 {
513    100
514}
515
516/// Default burst capacity above the sustained rate: 20 requests.
517///
518/// **Why 20?** Allows a short burst (e.g. a web page loading many parallel
519/// requests) without immediately triggering the rate limiter, while still
520/// preventing sustained abuse.  Equivalent to ~200 ms of sustained quota at
521/// the 100 req/s default.
522fn default_burst() -> u32 {
523    20
524}
525
526/// Default enabled state for rate limiting: `false`.
527///
528/// **Why false?** Rate limiting has no sensible universal default — the right
529/// limit depends entirely on the provider contract.  Enabling it by default
530/// with wrong values would silently reject legitimate traffic.  Operators must
531/// explicitly set `enabled = true` and configure `requests_per_second`.
532fn default_rate_limit_enabled() -> bool {
533    false
534}
535
536/// Rate limiting configuration.
537///
538/// Controls the token-bucket rate limiter that caps inbound request throughput.
539/// When enabled, requests exceeding the allowed rate are rejected with HTTP 429.
540///
541/// # Panics
542///
543/// This type never panics.
544#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
545#[cfg_attr(feature = "schema", derive(JsonSchema))]
546#[serde(deny_unknown_fields)]
547pub struct RateLimitConfig {
548    /// Whether rate limiting is enabled.
549    #[serde(default = "default_rate_limit_enabled")]
550    pub enabled: bool,
551    /// Maximum sustained requests per second (token refill rate).
552    #[serde(default = "default_rps")]
553    pub requests_per_second: u32,
554    /// Maximum burst capacity above the sustained rate.
555    #[serde(default = "default_burst")]
556    pub burst_capacity: u32,
557}
558
559impl Default for RateLimitConfig {
560    fn default() -> Self {
561        Self {
562            enabled: default_rate_limit_enabled(),
563            requests_per_second: default_rps(),
564            burst_capacity: default_burst(),
565        }
566    }
567}
568
569// ── Resilience ───────────────────────────────────────────────────────────
570
571/// Resilience configuration for retries and circuit breaking.
572///
573/// # Panics
574///
575/// This type never panics.
576#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
577#[cfg_attr(feature = "schema", derive(JsonSchema))]
578#[serde(deny_unknown_fields)]
579pub struct ResilienceConfig {
580    /// Maximum number of retry attempts before failing a request.
581    pub retry_attempts: u32,
582    /// Base delay (ms) for exponential backoff. Must be ≤ `retry_max_ms`.
583    #[serde(default = "default_retry_base_ms")]
584    pub retry_base_ms: u64,
585    /// Maximum delay (ms) cap for exponential backoff.
586    #[serde(default = "default_retry_max_ms")]
587    pub retry_max_ms: u64,
588    /// Number of consecutive failures before the circuit breaker opens.
589    pub circuit_breaker_threshold: u32,
590    /// Seconds to keep the circuit breaker open before allowing a probe.
591    pub circuit_breaker_timeout_s: u64,
592    /// Required success rate (0.0–1.0) over the sliding window.
593    pub circuit_breaker_success_rate: f64,
594}
595
596// ── Deduplication ────────────────────────────────────────────────────────
597
598/// Deduplication configuration for request coalescing.
599///
600/// # Panics
601///
602/// This type never panics.
603#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
604#[cfg_attr(feature = "schema", derive(JsonSchema))]
605#[serde(deny_unknown_fields)]
606pub struct DeduplicationConfig {
607    /// Whether deduplication is enabled.
608    pub enabled: bool,
609    /// Time window (seconds) within which duplicate requests are coalesced.
610    #[serde(default = "default_dedup_window_s")]
611    pub window_s: u64,
612    /// Maximum number of dedup entries held in memory.
613    #[serde(default = "default_dedup_max_entries")]
614    pub max_entries: usize,
615}
616
617// ── Observability ────────────────────────────────────────────────────────
618
619/// Observability configuration: logging, metrics endpoint, and tracing.
620///
621/// # Panics
622///
623/// This type never panics.
624#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
625#[cfg_attr(feature = "schema", derive(JsonSchema))]
626#[serde(deny_unknown_fields)]
627pub struct ObservabilityConfig {
628    /// Log output format.
629    pub log_format: LogFormat,
630    /// Port for the Prometheus metrics HTTP endpoint. `None` disables it.
631    pub metrics_port: Option<u16>,
632    /// OpenTelemetry tracing collector endpoint. `None` disables distributed tracing.
633    pub tracing_endpoint: Option<String>,
634}
635
636/// Log output format.
637///
638/// # Panics
639///
640/// This type never panics.
641#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
642#[cfg_attr(feature = "schema", derive(JsonSchema))]
643#[serde(deny_unknown_fields)]
644#[serde(rename_all = "snake_case")]
645pub enum LogFormat {
646    /// Human-readable, colorized log output.
647    Pretty,
648    /// Structured JSON log output for machine consumption.
649    Json,
650}
651
652/// Export the JSON Schema for `PipelineConfig`.
653///
654/// This enables IDE autocomplete when editing TOML config files.
655///
656/// # Errors
657///
658/// Returns `serde_json::Error` if schema serialization fails (should not
659/// happen with well-formed derive macros).
660///
661/// # Panics
662///
663/// This function never panics.
664#[cfg(feature = "schema")]
665pub fn export_schema() -> Result<String, serde_json::Error> {
666    let schema = schemars::schema_for!(PipelineConfig);
667    serde_json::to_string_pretty(&schema)
668}
669
670#[cfg(test)]
671mod tests {
672    use super::*;
673
674    #[test]
675    fn test_default_timeout_ms_returns_5000() {
676        assert_eq!(default_timeout_ms(), 5000);
677    }
678
679    #[test]
680    fn test_default_max_context_tokens_returns_2048() {
681        assert_eq!(default_max_context_tokens(), 2048);
682    }
683
684    #[test]
685    fn test_default_retry_base_ms_returns_100() {
686        assert_eq!(default_retry_base_ms(), 100);
687    }
688
689    #[test]
690    fn test_default_retry_max_ms_returns_5000() {
691        assert_eq!(default_retry_max_ms(), 5000);
692    }
693
694    #[test]
695    fn test_default_dedup_window_s_returns_300() {
696        assert_eq!(default_dedup_window_s(), 300);
697    }
698
699    #[test]
700    fn test_default_dedup_max_entries_returns_10000() {
701        assert_eq!(default_dedup_max_entries(), 10_000);
702    }
703
704    #[test]
705    fn test_default_channel_capacity_returns_512() {
706        assert_eq!(default_channel_capacity(), 512);
707    }
708
709    #[test]
710    fn test_default_true_returns_true() {
711        assert!(default_true());
712    }
713
714    #[test]
715    fn test_worker_kind_serializes_to_snake_case() {
716        let json = serde_json::to_string(&WorkerKind::OpenAi).expect("test: serialization");
717        assert_eq!(json, "\"open_ai\"");
718    }
719
720    #[test]
721    fn test_worker_kind_deserializes_from_snake_case() {
722        let kind: WorkerKind =
723            serde_json::from_str("\"llama_cpp\"").expect("test: deserialization");
724        assert_eq!(kind, WorkerKind::LlamaCpp);
725    }
726
727    #[test]
728    fn test_log_format_serializes_to_snake_case() {
729        let json = serde_json::to_string(&LogFormat::Pretty).expect("test: serialization");
730        assert_eq!(json, "\"pretty\"");
731    }
732
733    #[test]
734    fn test_log_format_deserializes_from_snake_case() {
735        let fmt: LogFormat = serde_json::from_str("\"json\"").expect("test: deserialization");
736        assert_eq!(fmt, LogFormat::Json);
737    }
738
739    #[cfg(feature = "schema")]
740    #[test]
741    fn test_export_schema_produces_valid_json() {
742        let schema = export_schema().expect("test: schema export");
743        let parsed: serde_json::Value =
744            serde_json::from_str(&schema).expect("test: schema is valid JSON");
745        // Should contain top-level properties
746        assert!(parsed.get("properties").is_some() || parsed.get("$ref").is_some());
747    }
748
749    #[test]
750    fn test_pipeline_config_minimal_toml_parses() {
751        let toml_str = r#"
752[pipeline]
753name = "test"
754version = "1.0"
755
756[stages.rag]
757enabled = true
758
759[stages.assemble]
760enabled = true
761
762[stages.inference]
763worker = "echo"
764model = "test-model"
765
766[stages.post_process]
767enabled = true
768
769[stages.stream]
770enabled = true
771
772[resilience]
773retry_attempts = 3
774circuit_breaker_threshold = 5
775circuit_breaker_timeout_s = 60
776circuit_breaker_success_rate = 0.8
777
778[deduplication]
779enabled = false
780
781[observability]
782log_format = "pretty"
783"#;
784        let config: PipelineConfig = toml::from_str(toml_str).expect("test: minimal TOML parses");
785        assert_eq!(config.pipeline.name, "test");
786        assert_eq!(config.stages.inference.worker, WorkerKind::Echo);
787        assert_eq!(config.resilience.retry_base_ms, 100); // default applied
788        assert!(!config.deduplication.enabled);
789    }
790
791    #[test]
792    fn test_pipeline_config_full_toml_parses() {
793        let toml_str = r#"
794[pipeline]
795name = "production"
796version = "1.0"
797description = "Production pipeline with OpenAI GPT-4"
798
799[stages.rag]
800enabled = true
801timeout_ms = 5000
802max_context_tokens = 2048
803
804[stages.assemble]
805enabled = true
806channel_capacity = 256
807
808[stages.inference]
809worker = "open_ai"
810model = "gpt-4"
811max_tokens = 1024
812temperature = 0.7
813timeout_ms = 30000
814
815[stages.post_process]
816enabled = true
817
818[stages.stream]
819enabled = true
820
821[resilience]
822retry_attempts = 3
823retry_base_ms = 100
824retry_max_ms = 5000
825circuit_breaker_threshold = 5
826circuit_breaker_timeout_s = 60
827circuit_breaker_success_rate = 0.8
828
829[deduplication]
830enabled = true
831window_s = 300
832max_entries = 10000
833
834[observability]
835log_format = "json"
836metrics_port = 9090
837"#;
838        let config: PipelineConfig = toml::from_str(toml_str).expect("test: full TOML parses");
839        assert_eq!(config.pipeline.name, "production");
840        assert_eq!(config.stages.inference.worker, WorkerKind::OpenAi);
841        assert_eq!(config.stages.inference.temperature, Some(0.7));
842        assert_eq!(config.stages.assemble.channel_capacity, 256);
843        assert_eq!(config.observability.metrics_port, Some(9090));
844    }
845
846    #[test]
847    fn test_pipeline_config_serialize_deserialize_roundtrip() {
848        let config = PipelineConfig {
849            pipeline: PipelineSection {
850                name: "roundtrip".into(),
851                version: "2.0".into(),
852                description: Some("Roundtrip test".into()),
853            },
854            stages: StagesConfig {
855                rag: RagStageConfig {
856                    enabled: true,
857                    timeout_ms: 3000,
858                    max_context_tokens: 1024,
859                    channel_capacity: Some(256),
860                },
861                assemble: AssembleStageConfig {
862                    enabled: true,
863                    channel_capacity: 512,
864                },
865                inference: InferenceStageConfig {
866                    worker: WorkerKind::Anthropic,
867                    model: "claude-3-opus".into(),
868                    max_tokens: Some(2048),
869                    temperature: Some(0.5),
870                    timeout_ms: Some(60000),
871                    inference_workers: 1,
872                    adaptive_timeout_enabled: false,
873                    adaptive_timeout_min_ms: 500,
874                },
875                post_process: PostProcessStageConfig {
876                    enabled: true,
877                    channel_capacity: 512,
878                },
879                stream: StreamStageConfig {
880                    enabled: false,
881                    channel_capacity: 128,
882                },
883            },
884            resilience: ResilienceConfig {
885                retry_attempts: 5,
886                retry_base_ms: 200,
887                retry_max_ms: 10000,
888                circuit_breaker_threshold: 10,
889                circuit_breaker_timeout_s: 120,
890                circuit_breaker_success_rate: 0.9,
891            },
892            deduplication: DeduplicationConfig {
893                enabled: true,
894                window_s: 600,
895                max_entries: 20000,
896            },
897            observability: ObservabilityConfig {
898                log_format: LogFormat::Json,
899                metrics_port: Some(8080),
900                tracing_endpoint: Some("http://jaeger:14268".into()),
901            },
902            rate_limits: RateLimitConfig::default(),
903            channel_sizes: None,
904            workers: WorkersConfig::default(),
905        };
906
907        let toml_str = toml::to_string_pretty(&config).expect("test: serialize to TOML");
908        let deserialized: PipelineConfig =
909            toml::from_str(&toml_str).expect("test: deserialize from TOML");
910        assert_eq!(config, deserialized);
911    }
912
913    #[test]
914    fn test_pipeline_config_json_roundtrip() {
915        let config = PipelineConfig {
916            pipeline: PipelineSection {
917                name: "json-rt".into(),
918                version: "1.0".into(),
919                description: None,
920            },
921            stages: StagesConfig {
922                rag: RagStageConfig {
923                    enabled: true,
924                    timeout_ms: 5000,
925                    max_context_tokens: 2048,
926                    channel_capacity: None,
927                },
928                assemble: AssembleStageConfig {
929                    enabled: true,
930                    channel_capacity: 512,
931                },
932                inference: InferenceStageConfig {
933                    worker: WorkerKind::Echo,
934                    model: "echo".into(),
935                    max_tokens: None,
936                    temperature: None,
937                    timeout_ms: None,
938                    inference_workers: 1,
939                    adaptive_timeout_enabled: false,
940                    adaptive_timeout_min_ms: 500,
941                },
942                post_process: PostProcessStageConfig {
943                    enabled: true,
944                    channel_capacity: 512,
945                },
946                stream: StreamStageConfig {
947                    enabled: true,
948                    channel_capacity: 512,
949                },
950            },
951            resilience: ResilienceConfig {
952                retry_attempts: 3,
953                retry_base_ms: 100,
954                retry_max_ms: 5000,
955                circuit_breaker_threshold: 5,
956                circuit_breaker_timeout_s: 60,
957                circuit_breaker_success_rate: 0.8,
958            },
959            deduplication: DeduplicationConfig {
960                enabled: false,
961                window_s: 300,
962                max_entries: 10000,
963            },
964            observability: ObservabilityConfig {
965                log_format: LogFormat::Pretty,
966                metrics_port: None,
967                tracing_endpoint: None,
968            },
969            rate_limits: RateLimitConfig::default(),
970            channel_sizes: None,
971            workers: WorkersConfig::default(),
972        };
973
974        let json = serde_json::to_string(&config).expect("test: serialize to JSON");
975        let deserialized: PipelineConfig =
976            serde_json::from_str(&json).expect("test: deserialize from JSON");
977        assert_eq!(config, deserialized);
978    }
979
980    #[test]
981    fn test_all_worker_kinds_roundtrip_toml() {
982        // TOML requires a table wrapper for enum serialization
983        #[derive(Debug, Serialize, Deserialize, PartialEq)]
984        struct Wrapper {
985            kind: WorkerKind,
986        }
987
988        let kinds = vec![
989            WorkerKind::OpenAi,
990            WorkerKind::Anthropic,
991            WorkerKind::LlamaCpp,
992            WorkerKind::Vllm,
993            WorkerKind::Echo,
994        ];
995        for kind in kinds {
996            let w = Wrapper { kind: kind.clone() };
997            let s = toml::to_string(&w).expect("test: serialize worker kind");
998            let deserialized: Wrapper = toml::from_str(&s).expect("test: deserialize worker kind");
999            assert_eq!(w, deserialized);
1000        }
1001    }
1002
1003    #[test]
1004    fn test_pipeline_section_optional_description_omitted() {
1005        let toml_str = r#"
1006name = "no-desc"
1007version = "1.0"
1008"#;
1009        let section: PipelineSection =
1010            toml::from_str(toml_str).expect("test: parse without description");
1011        assert!(section.description.is_none());
1012    }
1013
1014    #[test]
1015    fn test_rag_stage_defaults_applied_when_omitted() {
1016        let toml_str = r#"
1017enabled = true
1018"#;
1019        let rag: RagStageConfig = toml::from_str(toml_str).expect("test: parse with defaults");
1020        assert_eq!(rag.timeout_ms, 5000);
1021        assert_eq!(rag.max_context_tokens, 2048);
1022        assert!(rag.channel_capacity.is_none());
1023    }
1024
1025    #[test]
1026    fn test_resilience_defaults_applied_when_omitted() {
1027        let toml_str = r#"
1028retry_attempts = 3
1029circuit_breaker_threshold = 5
1030circuit_breaker_timeout_s = 60
1031circuit_breaker_success_rate = 0.8
1032"#;
1033        let resilience: ResilienceConfig =
1034            toml::from_str(toml_str).expect("test: parse with defaults");
1035        assert_eq!(resilience.retry_base_ms, 100);
1036        assert_eq!(resilience.retry_max_ms, 5000);
1037    }
1038
1039    #[test]
1040    fn test_deduplication_defaults_applied_when_omitted() {
1041        let toml_str = r#"
1042enabled = true
1043"#;
1044        let dedup: DeduplicationConfig =
1045            toml::from_str(toml_str).expect("test: parse with defaults");
1046        assert_eq!(dedup.window_s, 300);
1047        assert_eq!(dedup.max_entries, 10_000);
1048    }
1049
1050    #[test]
1051    fn test_unknown_field_in_resilience_is_rejected() {
1052        // deny_unknown_fields must fire so typos in TOML are caught at parse time
1053        let bad_toml = r#"
1054retry_attempts = 3
1055circuit_breaker_threshold = 5
1056circuit_breaker_timeout_s = 60
1057circuit_breaker_success_rate = 0.8
1058typo_field = "oops"
1059"#;
1060        let result = toml::from_str::<ResilienceConfig>(bad_toml);
1061        assert!(result.is_err(), "unknown fields must be rejected");
1062        let msg = result.unwrap_err().to_string();
1063        assert!(
1064            msg.contains("typo_field") || msg.contains("unknown"),
1065            "error should mention the unknown field: {msg}"
1066        );
1067    }
1068
1069    #[test]
1070    fn test_unknown_field_in_deduplication_is_rejected() {
1071        let bad_toml = r#"
1072enabled = true
1073unknwon_key = 42
1074"#;
1075        let result = toml::from_str::<DeduplicationConfig>(bad_toml);
1076        assert!(result.is_err(), "unknown fields must be rejected");
1077    }
1078}