1use std::path::PathBuf;
2
3use serde::Deserialize;
4
5#[derive(Debug, Deserialize)]
7#[non_exhaustive]
8pub struct ServerConfig {
9 #[serde(default = "default_listen_addr")]
11 pub listen_addr: String,
12 #[serde(default = "default_listen_port")]
14 pub listen_port: u16,
15 pub tls_cert_path: Option<PathBuf>,
17 pub tls_key_path: Option<PathBuf>,
19 #[serde(default = "default_tls_handshake_timeout")]
24 pub tls_handshake_timeout: String,
25 #[serde(default = "default_max_concurrent_tls_handshakes")]
30 pub max_concurrent_tls_handshakes: usize,
31 #[serde(default = "default_shutdown_timeout")]
33 pub shutdown_timeout: String,
34 #[serde(default = "default_request_timeout")]
36 pub request_timeout: String,
37 #[serde(default)]
41 pub allowed_origins: Vec<String>,
42 #[serde(default)]
45 pub stdio_enabled: bool,
46 pub tool_rate_limit: Option<u32>,
50 pub tool_rate_limit_burst: Option<u32>,
54 pub extra_route_rate_limit: Option<u32>,
60 pub extra_route_rate_limit_burst: Option<u32>,
64 #[serde(default)]
70 pub extra_route_rate_limit_exempt_paths: Vec<String>,
71 #[serde(default)]
77 pub trusted_proxies: Vec<String>,
78 pub forwarded_header: Option<crate::transport::ForwardedHeaderMode>,
82 #[serde(default = "default_session_idle_timeout")]
85 pub session_idle_timeout: String,
86 #[serde(default = "default_sse_keep_alive")]
90 pub sse_keep_alive: String,
91 pub public_url: Option<String>,
96 #[serde(default)]
98 pub compression_enabled: bool,
99 #[serde(default = "default_compression_min_size")]
102 pub compression_min_size: u16,
103 pub max_concurrent_requests: Option<usize>,
106 #[serde(default)]
108 pub admin_enabled: bool,
109 #[serde(default = "default_admin_role")]
111 pub admin_role: String,
112 pub auth: Option<crate::auth::AuthConfig>,
114}
115
116impl Default for ServerConfig {
117 fn default() -> Self {
118 Self {
119 listen_addr: default_listen_addr(),
120 listen_port: default_listen_port(),
121 tls_cert_path: None,
122 tls_key_path: None,
123 tls_handshake_timeout: default_tls_handshake_timeout(),
124 max_concurrent_tls_handshakes: default_max_concurrent_tls_handshakes(),
125 shutdown_timeout: default_shutdown_timeout(),
126 request_timeout: default_request_timeout(),
127 allowed_origins: Vec::new(),
128 stdio_enabled: false,
129 tool_rate_limit: None,
130 tool_rate_limit_burst: None,
131 extra_route_rate_limit: None,
132 extra_route_rate_limit_burst: None,
133 extra_route_rate_limit_exempt_paths: Vec::new(),
134 trusted_proxies: Vec::new(),
135 forwarded_header: None,
136 session_idle_timeout: default_session_idle_timeout(),
137 sse_keep_alive: default_sse_keep_alive(),
138 public_url: None,
139 compression_enabled: false,
140 compression_min_size: default_compression_min_size(),
141 max_concurrent_requests: None,
142 admin_enabled: false,
143 admin_role: default_admin_role(),
144 auth: None,
145 }
146 }
147}
148
149#[derive(Debug, Deserialize)]
151#[non_exhaustive]
152pub struct ObservabilityConfig {
153 #[serde(default = "default_log_level")]
155 pub log_level: String,
156 #[serde(default = "default_log_format")]
158 pub log_format: String,
159 pub audit_log_path: Option<PathBuf>,
161 #[serde(default)]
164 pub log_request_headers: bool,
165 #[serde(default)]
167 pub metrics_enabled: bool,
168 #[serde(default = "default_metrics_bind")]
170 pub metrics_bind: String,
171}
172
173impl Default for ObservabilityConfig {
174 fn default() -> Self {
175 Self {
176 log_level: default_log_level(),
177 log_format: default_log_format(),
178 audit_log_path: None,
179 log_request_headers: false,
180 metrics_enabled: false,
181 metrics_bind: default_metrics_bind(),
182 }
183 }
184}
185
186pub fn validate_server_config(server: &ServerConfig) -> crate::error::Result<()> {
192 use crate::error::McpxError;
193
194 if server.listen_port == 0 {
195 return Err(McpxError::Config("listen_port must be nonzero".into()));
196 }
197
198 match (&server.tls_cert_path, &server.tls_key_path) {
199 (Some(_), None) | (None, Some(_)) => {
200 return Err(McpxError::Config(
201 "tls_cert_path and tls_key_path must both be set or both omitted".into(),
202 ));
203 }
204 _ => {}
205 }
206
207 if server.max_concurrent_requests == Some(0) {
208 return Err(McpxError::Config(
209 "max_concurrent_requests must be nonzero when set".into(),
210 ));
211 }
212
213 if server.extra_route_rate_limit == Some(0) {
214 return Err(McpxError::Config(
215 "server.extra_route_rate_limit must be greater than zero".into(),
216 ));
217 }
218
219 validate_rate_limit_knobs(server)?;
220 validate_trusted_forwarder_config(server)?;
221
222 if server.admin_enabled {
223 let auth_enabled = server.auth.as_ref().is_some_and(|a| a.enabled);
224 if !auth_enabled {
225 return Err(McpxError::Config(
226 "admin_enabled=true requires auth to be configured and enabled".into(),
227 ));
228 }
229 if server.admin_role.trim().is_empty() {
230 return Err(McpxError::Config("admin_role must not be empty".into()));
231 }
232 }
233
234 for (field, value) in [
235 ("server.shutdown_timeout", server.shutdown_timeout.as_str()),
236 ("server.request_timeout", server.request_timeout.as_str()),
237 (
238 "server.session_idle_timeout",
239 server.session_idle_timeout.as_str(),
240 ),
241 ("server.sse_keep_alive", server.sse_keep_alive.as_str()),
242 (
243 "server.tls_handshake_timeout",
244 server.tls_handshake_timeout.as_str(),
245 ),
246 ] {
247 if humantime::parse_duration(value).is_err() {
248 return Err(McpxError::Config(format!(
249 "invalid duration for {field}: {value:?}"
250 )));
251 }
252 }
253
254 if humantime::parse_duration(&server.tls_handshake_timeout)
258 .is_ok_and(|d| d == std::time::Duration::ZERO)
259 {
260 return Err(McpxError::Config(
261 "server.tls_handshake_timeout must be greater than zero".into(),
262 ));
263 }
264
265 if server.max_concurrent_tls_handshakes == 0 {
269 return Err(McpxError::Config(
270 "server.max_concurrent_tls_handshakes must be greater than zero".into(),
271 ));
272 }
273
274 Ok(())
275}
276
277fn validate_rate_limit_knobs(server: &ServerConfig) -> crate::error::Result<()> {
281 use crate::error::McpxError;
282
283 if server.tool_rate_limit_burst == Some(0) {
284 return Err(McpxError::Config(
285 "server.tool_rate_limit_burst must be greater than zero".into(),
286 ));
287 }
288 if server.extra_route_rate_limit_burst == Some(0) {
289 return Err(McpxError::Config(
290 "server.extra_route_rate_limit_burst must be greater than zero".into(),
291 ));
292 }
293 if server.tool_rate_limit_burst.is_some() && server.tool_rate_limit.is_none() {
294 return Err(McpxError::Config(
295 "server.tool_rate_limit_burst requires server.tool_rate_limit".into(),
296 ));
297 }
298 if server.extra_route_rate_limit_burst.is_some() && server.extra_route_rate_limit.is_none() {
299 return Err(McpxError::Config(
300 "server.extra_route_rate_limit_burst requires server.extra_route_rate_limit".into(),
301 ));
302 }
303 if !server.extra_route_rate_limit_exempt_paths.is_empty()
304 && server.extra_route_rate_limit.is_none()
305 {
306 return Err(McpxError::Config(
307 "server.extra_route_rate_limit_exempt_paths requires server.extra_route_rate_limit"
308 .into(),
309 ));
310 }
311 for path in &server.extra_route_rate_limit_exempt_paths {
312 if path.is_empty() || !path.starts_with('/') {
313 return Err(McpxError::Config(format!(
314 "server.extra_route_rate_limit_exempt_paths entries must be non-empty and start with '/': {path:?}"
315 )));
316 }
317 }
318 if let Some(rl) = server.auth.as_ref().and_then(|a| a.rate_limit.as_ref()) {
319 if rl.burst == Some(0) {
320 return Err(McpxError::Config(
321 "auth.rate_limit.burst must be greater than zero".into(),
322 ));
323 }
324 if rl.pre_auth_burst == Some(0) {
325 return Err(McpxError::Config(
326 "auth.rate_limit.pre_auth_burst must be greater than zero".into(),
327 ));
328 }
329 }
330 Ok(())
331}
332
333fn validate_trusted_forwarder_config(server: &ServerConfig) -> crate::error::Result<()> {
336 use crate::error::McpxError;
337
338 for entry in &server.trusted_proxies {
339 crate::transport::validate_trusted_proxy_entry(entry).map_err(McpxError::Config)?;
340 }
341 if server.forwarded_header.is_some() && server.trusted_proxies.is_empty() {
342 return Err(McpxError::Config(
343 "server.forwarded_header requires server.trusted_proxies to be nonempty".into(),
344 ));
345 }
346 Ok(())
347}
348
349pub fn validate_observability_config(obs: &ObservabilityConfig) -> crate::error::Result<()> {
355 use tracing_subscriber::EnvFilter;
356
357 use crate::error::McpxError;
358
359 if EnvFilter::try_new(&obs.log_level).is_err() {
360 return Err(McpxError::Config(format!(
361 "invalid log_level: {:?} (expected a valid tracing filter directive, e.g. \"info\", \"debug,hyper=warn\")",
362 obs.log_level
363 )));
364 }
365 let valid_formats = ["json", "pretty", "text"];
366 if !valid_formats.contains(&obs.log_format.as_str()) {
367 return Err(McpxError::Config(format!(
368 "invalid log_format: {:?} (expected one of: {valid_formats:?})",
369 obs.log_format
370 )));
371 }
372
373 Ok(())
374}
375
376fn default_listen_addr() -> String {
379 "127.0.0.1".into()
380}
381fn default_listen_port() -> u16 {
382 8443
383}
384fn default_shutdown_timeout() -> String {
385 "30s".into()
386}
387fn default_request_timeout() -> String {
388 "120s".into()
389}
390fn default_log_level() -> String {
391 "info,rmcp=warn".into()
392}
393fn default_log_format() -> String {
394 "pretty".into()
395}
396fn default_metrics_bind() -> String {
397 "127.0.0.1:9090".into()
398}
399fn default_session_idle_timeout() -> String {
400 "20m".into()
401}
402fn default_tls_handshake_timeout() -> String {
403 "10s".into()
404}
405const fn default_max_concurrent_tls_handshakes() -> usize {
406 256
407}
408fn default_admin_role() -> String {
409 "admin".into()
410}
411fn default_compression_min_size() -> u16 {
412 1024
413}
414fn default_sse_keep_alive() -> String {
415 "15s".into()
416}
417
418#[cfg(test)]
419mod tests {
420 #![allow(
421 clippy::unwrap_used,
422 clippy::expect_used,
423 clippy::panic,
424 clippy::indexing_slicing,
425 clippy::unwrap_in_result,
426 clippy::print_stdout,
427 clippy::print_stderr,
428 reason = "test-only relaxations; production code uses ? and tracing"
429 )]
430 use super::*;
431
432 #[test]
435 fn server_config_defaults() {
436 let cfg = ServerConfig::default();
437 assert_eq!(cfg.listen_addr, "127.0.0.1");
438 assert_eq!(cfg.listen_port, 8443);
439 assert!(cfg.tls_cert_path.is_none());
440 assert!(cfg.tls_key_path.is_none());
441 assert_eq!(cfg.shutdown_timeout, "30s");
442 assert_eq!(cfg.request_timeout, "120s");
443 assert!(cfg.allowed_origins.is_empty());
444 assert!(!cfg.stdio_enabled);
445 assert!(cfg.tool_rate_limit.is_none());
446 assert_eq!(cfg.session_idle_timeout, "20m");
447 assert_eq!(cfg.sse_keep_alive, "15s");
448 assert!(cfg.public_url.is_none());
449 }
450
451 #[test]
452 fn observability_config_defaults() {
453 let cfg = ObservabilityConfig::default();
454 assert_eq!(cfg.log_level, "info,rmcp=warn");
455 assert_eq!(cfg.log_format, "pretty");
456 assert!(cfg.audit_log_path.is_none());
457 assert!(!cfg.log_request_headers);
458 assert!(!cfg.metrics_enabled);
459 assert_eq!(cfg.metrics_bind, "127.0.0.1:9090");
460 }
461
462 #[test]
465 fn valid_server_config_passes() {
466 let cfg = ServerConfig::default();
467 assert!(validate_server_config(&cfg).is_ok());
468 }
469
470 #[test]
471 fn zero_port_rejected() {
472 let cfg = ServerConfig {
473 listen_port: 0,
474 ..ServerConfig::default()
475 };
476 let err = validate_server_config(&cfg).unwrap_err();
477 assert!(err.to_string().contains("listen_port"));
478 }
479
480 #[test]
481 fn zero_extra_route_rate_limit_rejected() {
482 let cfg = ServerConfig {
483 extra_route_rate_limit: Some(0),
484 ..ServerConfig::default()
485 };
486 let err = validate_server_config(&cfg).unwrap_err();
487 assert!(err.to_string().contains("extra_route_rate_limit"));
488 }
489
490 #[test]
491 fn zero_burst_knobs_rejected() {
492 let cfg = ServerConfig {
493 tool_rate_limit: Some(10),
494 tool_rate_limit_burst: Some(0),
495 ..ServerConfig::default()
496 };
497 let err = validate_server_config(&cfg).unwrap_err();
498 assert!(err.to_string().contains("tool_rate_limit_burst"));
499
500 let cfg = ServerConfig {
501 extra_route_rate_limit: Some(10),
502 extra_route_rate_limit_burst: Some(0),
503 ..ServerConfig::default()
504 };
505 let err = validate_server_config(&cfg).unwrap_err();
506 assert!(err.to_string().contains("extra_route_rate_limit_burst"));
507 }
508
509 #[test]
510 fn orphan_burst_knobs_rejected() {
511 let cfg = ServerConfig {
512 tool_rate_limit_burst: Some(5),
513 ..ServerConfig::default()
514 };
515 let err = validate_server_config(&cfg).unwrap_err();
516 assert!(err.to_string().contains("requires server.tool_rate_limit"));
517
518 let cfg = ServerConfig {
519 extra_route_rate_limit_burst: Some(5),
520 ..ServerConfig::default()
521 };
522 let err = validate_server_config(&cfg).unwrap_err();
523 assert!(
524 err.to_string()
525 .contains("requires server.extra_route_rate_limit")
526 );
527 }
528
529 #[test]
530 fn exempt_paths_toml_roundtrip_and_validation() {
531 let cfg: ServerConfig = toml::from_str(
532 r#"
533 extra_route_rate_limit = 60
534 extra_route_rate_limit_exempt_paths = ["/.well-known/oauth-authorization-server"]
535 "#,
536 )
537 .unwrap();
538 assert_eq!(
539 cfg.extra_route_rate_limit_exempt_paths,
540 vec!["/.well-known/oauth-authorization-server".to_owned()]
541 );
542 assert!(validate_server_config(&cfg).is_ok());
543 }
544
545 #[test]
546 fn orphan_exempt_paths_rejected() {
547 let cfg = ServerConfig {
548 extra_route_rate_limit_exempt_paths: vec!["/ok".into()],
549 ..ServerConfig::default()
550 };
551 let err = validate_server_config(&cfg).unwrap_err();
552 assert!(
553 err.to_string()
554 .contains("requires server.extra_route_rate_limit")
555 );
556 }
557
558 #[test]
559 fn malformed_exempt_paths_rejected() {
560 for bad in ["", "no-slash"] {
561 let cfg = ServerConfig {
562 extra_route_rate_limit: Some(10),
563 extra_route_rate_limit_exempt_paths: vec![bad.into()],
564 ..ServerConfig::default()
565 };
566 let err = validate_server_config(&cfg).unwrap_err();
567 assert!(
568 err.to_string()
569 .contains("must be non-empty and start with '/'"),
570 "entry {bad:?}: {err}"
571 );
572 }
573 }
574
575 #[test]
576 fn bad_trusted_proxy_entry_rejected() {
577 let cfg = ServerConfig {
578 trusted_proxies: vec!["not-a-cidr".into()],
579 ..ServerConfig::default()
580 };
581 let err = validate_server_config(&cfg).unwrap_err();
582 assert!(err.to_string().contains("trusted_proxies"));
583 }
584
585 #[test]
586 fn zero_prefix_trusted_proxy_rejected() {
587 for entry in ["0.0.0.0/0", "::/0"] {
588 let cfg = ServerConfig {
589 trusted_proxies: vec![entry.into()],
590 ..ServerConfig::default()
591 };
592 let err = validate_server_config(&cfg).unwrap_err();
593 assert!(
594 err.to_string().contains("prefix length 0"),
595 "entry {entry:?}: {err}"
596 );
597 }
598 }
599
600 #[test]
601 fn cidr_and_bare_ip_proxy_entries_accepted() {
602 let cfg = ServerConfig {
603 trusted_proxies: vec!["10.0.0.0/8".into(), "192.0.2.1".into()],
604 ..ServerConfig::default()
605 };
606 assert!(validate_server_config(&cfg).is_ok());
607 }
608
609 #[test]
610 fn forwarded_header_without_proxies_rejected() {
611 let cfg = ServerConfig {
612 forwarded_header: Some(crate::transport::ForwardedHeaderMode::Forwarded),
613 ..ServerConfig::default()
614 };
615 let err = validate_server_config(&cfg).unwrap_err();
616 assert!(err.to_string().contains("requires server.trusted_proxies"));
617 }
618
619 #[test]
620 fn zero_auth_bursts_rejected() {
621 let auth = crate::auth::AuthConfig::with_keys(vec![])
622 .with_rate_limit(crate::auth::RateLimitConfig::new(10).with_burst(0));
623 let cfg = ServerConfig {
624 auth: Some(auth),
625 ..ServerConfig::default()
626 };
627 let err = validate_server_config(&cfg).unwrap_err();
628 assert!(err.to_string().contains("rate_limit.burst"));
629
630 let auth = crate::auth::AuthConfig::with_keys(vec![])
631 .with_rate_limit(crate::auth::RateLimitConfig::new(10).with_pre_auth_burst(0));
632 let cfg = ServerConfig {
633 auth: Some(auth),
634 ..ServerConfig::default()
635 };
636 let err = validate_server_config(&cfg).unwrap_err();
637 assert!(err.to_string().contains("pre_auth_burst"));
638 }
639
640 #[test]
641 fn tls_cert_without_key_rejected() {
642 let cfg = ServerConfig {
643 tls_cert_path: Some("/tmp/cert.pem".into()),
644 ..ServerConfig::default()
645 };
646 let err = validate_server_config(&cfg).unwrap_err();
647 assert!(err.to_string().contains("tls_cert_path"));
648 }
649
650 #[test]
651 fn tls_key_without_cert_rejected() {
652 let cfg = ServerConfig {
653 tls_key_path: Some("/tmp/key.pem".into()),
654 ..ServerConfig::default()
655 };
656 let err = validate_server_config(&cfg).unwrap_err();
657 assert!(err.to_string().contains("tls_cert_path"));
658 }
659
660 #[test]
661 fn tls_both_set_passes() {
662 let cfg = ServerConfig {
663 tls_cert_path: Some("/tmp/cert.pem".into()),
664 tls_key_path: Some("/tmp/key.pem".into()),
665 ..ServerConfig::default()
666 };
667 assert!(validate_server_config(&cfg).is_ok());
668 }
669
670 #[test]
671 fn invalid_tls_handshake_timeout_rejected() {
672 let cfg = ServerConfig {
673 tls_handshake_timeout: "not-a-duration".into(),
674 ..ServerConfig::default()
675 };
676 let err = validate_server_config(&cfg).unwrap_err();
677 assert!(err.to_string().contains("tls_handshake_timeout"));
678 }
679
680 #[test]
681 fn zero_tls_handshake_timeout_rejected() {
682 let cfg = ServerConfig {
683 tls_handshake_timeout: "0s".into(),
684 ..ServerConfig::default()
685 };
686 let err = validate_server_config(&cfg).unwrap_err();
687 assert!(err.to_string().contains("tls_handshake_timeout"));
688 }
689
690 #[test]
691 fn zero_max_concurrent_tls_handshakes_rejected() {
692 let cfg = ServerConfig {
693 max_concurrent_tls_handshakes: 0,
694 ..ServerConfig::default()
695 };
696 let err = validate_server_config(&cfg).unwrap_err();
697 assert!(err.to_string().contains("max_concurrent_tls_handshakes"));
698 }
699
700 #[test]
701 fn invalid_shutdown_timeout_rejected() {
702 let cfg = ServerConfig {
703 shutdown_timeout: "not-a-duration".into(),
704 ..ServerConfig::default()
705 };
706 let err = validate_server_config(&cfg).unwrap_err();
707 assert!(err.to_string().contains("shutdown_timeout"));
708 }
709
710 #[test]
711 fn invalid_request_timeout_rejected() {
712 let cfg = ServerConfig {
713 request_timeout: "xyz".into(),
714 ..ServerConfig::default()
715 };
716 let err = validate_server_config(&cfg).unwrap_err();
717 assert!(err.to_string().contains("request_timeout"));
718 }
719
720 #[test]
723 fn valid_observability_config_passes() {
724 let cfg = ObservabilityConfig::default();
725 assert!(validate_observability_config(&cfg).is_ok());
726 }
727
728 #[test]
729 fn invalid_log_level_rejected() {
730 let cfg = ObservabilityConfig {
731 log_level: "[invalid".into(),
732 ..ObservabilityConfig::default()
733 };
734 let err = validate_observability_config(&cfg).unwrap_err();
735 assert!(err.to_string().contains("log_level"));
736 }
737
738 #[test]
739 fn invalid_log_format_rejected() {
740 let cfg = ObservabilityConfig {
741 log_format: "yaml".into(),
742 ..ObservabilityConfig::default()
743 };
744 let err = validate_observability_config(&cfg).unwrap_err();
745 assert!(err.to_string().contains("log_format"));
746 }
747
748 #[test]
749 fn all_valid_log_levels_accepted() {
750 for level in &[
751 "trace",
752 "debug",
753 "info",
754 "warn",
755 "error",
756 "info,rmcp=warn",
757 "debug,hyper=error",
758 ] {
759 let cfg = ObservabilityConfig {
760 log_level: (*level).into(),
761 ..ObservabilityConfig::default()
762 };
763 assert!(
764 validate_observability_config(&cfg).is_ok(),
765 "level {level} should be valid"
766 );
767 }
768 }
769
770 #[test]
771 fn all_log_formats_accepted() {
772 for fmt in &["json", "pretty", "text"] {
773 let cfg = ObservabilityConfig {
774 log_format: (*fmt).into(),
775 ..ObservabilityConfig::default()
776 };
777 assert!(
778 validate_observability_config(&cfg).is_ok(),
779 "format {fmt} should be valid"
780 );
781 }
782 }
783
784 #[test]
787 fn server_config_deserialize_defaults() {
788 let cfg: ServerConfig = toml::from_str("").unwrap();
789 assert_eq!(cfg.listen_port, 8443);
790 assert_eq!(cfg.listen_addr, "127.0.0.1");
791 assert_eq!(cfg.tls_handshake_timeout, "10s");
792 assert_eq!(cfg.max_concurrent_tls_handshakes, 256);
793 }
794
795 #[test]
796 fn observability_config_deserialize_defaults() {
797 let cfg: ObservabilityConfig = toml::from_str("").unwrap();
798 assert_eq!(cfg.log_level, "info,rmcp=warn");
799 assert_eq!(cfg.log_format, "pretty");
800 assert!(!cfg.log_request_headers);
801 assert!(!cfg.metrics_enabled);
802 }
803}