use crate::route::{RouteTable, Upstream};
use serde::Deserialize;
#[derive(Debug, Clone, Deserialize)]
pub struct ConfigToml {
#[serde(default = "default_health_interval")]
pub health_interval_secs: u64,
#[serde(default = "default_cooldown")]
pub default_cooldown_secs: u64,
#[serde(default = "default_rps")]
pub tenant_rps: u32,
pub tls: TlsToml,
pub jwt: JwtToml,
pub routes: Vec<RouteToml>,
}
fn default_health_interval() -> u64 {
5
}
fn default_cooldown() -> u64 {
30
}
fn default_rps() -> u32 {
1_000
}
#[derive(Debug, Clone, Deserialize)]
pub struct TlsToml {
pub cert_path: String,
pub key_path: String,
pub client_ca_path: String,
}
#[derive(Debug, Clone, Deserialize)]
pub struct JwtToml {
pub jwks_path: String,
#[serde(default)]
pub issuer: Option<String>,
#[serde(default)]
pub audience: Option<String>,
}
#[derive(Debug, Clone, Deserialize)]
pub struct RouteToml {
pub prefix: String,
pub upstream: String,
#[serde(default)]
pub cooldown_secs: Option<u64>,
}
#[derive(Debug, Clone, Deserialize, Default)]
#[non_exhaustive]
pub struct ConfigExt {
#[serde(default)]
pub limits: Limits,
#[serde(default)]
pub routes: Vec<RouteExt>,
}
#[derive(Debug, Clone, Deserialize)]
#[non_exhaustive]
#[serde(default)]
pub struct Limits {
pub max_request_body_bytes: u64,
pub request_body_timeout_secs: u64,
pub upstream_timeout_secs: u64,
pub tls_handshake_timeout_secs: u64,
pub first_request_timeout_secs: u64,
pub h2_max_concurrent_streams: u32,
pub shutdown_drain_secs: u64,
}
impl Default for Limits {
fn default() -> Self {
Self {
max_request_body_bytes: 8 * 1024 * 1024,
request_body_timeout_secs: 30,
upstream_timeout_secs: 30,
tls_handshake_timeout_secs: 10,
first_request_timeout_secs: 10,
h2_max_concurrent_streams: 64,
shutdown_drain_secs: 25,
}
}
}
#[derive(Debug, Clone, Deserialize)]
#[non_exhaustive]
pub struct RouteExt {
pub prefix: String,
#[serde(default)]
pub health_path: Option<String>,
#[serde(default)]
pub health_disabled: bool,
}
const MAX_SECS: u64 = 86_400;
const MAX_BODY: u64 = 1 << 30;
impl ConfigExt {
fn validate(&self, cfg: &ConfigToml) -> anyhow::Result<()> {
let l = &self.limits;
for (key, v) in [
("request_body_timeout_secs", l.request_body_timeout_secs),
("upstream_timeout_secs", l.upstream_timeout_secs),
("tls_handshake_timeout_secs", l.tls_handshake_timeout_secs),
("first_request_timeout_secs", l.first_request_timeout_secs),
("shutdown_drain_secs", l.shutdown_drain_secs),
] {
anyhow::ensure!(
(1..=MAX_SECS).contains(&v),
"limits.{key} must be between 1 and {MAX_SECS}"
);
}
anyhow::ensure!(
(1..=MAX_BODY).contains(&l.max_request_body_bytes),
"limits.max_request_body_bytes must be between 1 and {MAX_BODY}"
);
anyhow::ensure!(
l.h2_max_concurrent_streams > 0,
"limits.h2_max_concurrent_streams must be at least 1"
);
anyhow::ensure!(
cfg.health_interval_secs > 0,
"health_interval_secs must be at least 1"
);
for r in &self.routes {
anyhow::ensure!(
cfg.routes.iter().any(|c| c.prefix == r.prefix),
"route {:?}: health keys match no [[routes]] prefix",
r.prefix
);
if let Some(p) = &r.health_path {
anyhow::ensure!(
p.starts_with('/')
&& !p.contains(['?', '#'])
&& p.parse::<http::uri::PathAndQuery>().is_ok(),
"route {}: health_path {p:?} must start with '/' and contain no '?' or '#'",
r.prefix
);
}
}
Ok(())
}
}
pub fn parse_config(raw: &str) -> anyhow::Result<(ConfigToml, ConfigExt)> {
let cfg: ConfigToml = toml::from_str(raw)?;
let ext: ConfigExt = toml::from_str(raw)?;
ext.validate(&cfg)?;
Ok((cfg, ext))
}
pub fn build_table(cfg: &ConfigToml) -> anyhow::Result<RouteTable> {
build_table_ext(cfg, &ConfigExt::default())
}
pub fn build_table_ext(cfg: &ConfigToml, ext: &ConfigExt) -> anyhow::Result<RouteTable> {
ext.validate(cfg)?;
let default_cooldown = cfg.default_cooldown_secs;
let mut rules = Vec::with_capacity(cfg.routes.len());
for r in &cfg.routes {
let uri: http::Uri = r.upstream.parse()?;
if uri.authority().is_none() {
anyhow::bail!(
"route {:?}: upstream {:?} has no authority (scheme://host)",
r.prefix,
r.upstream
);
}
let cooldown = r.cooldown_secs.unwrap_or(default_cooldown);
anyhow::ensure!(
cooldown > 0,
"route {}: cooldown_secs must be at least 1",
r.prefix
);
let mut up = Upstream::new(uri, cooldown);
if let Some(x) = ext.routes.iter().find(|x| x.prefix == r.prefix) {
up = up.with_health(x.health_path.clone(), x.health_disabled);
}
rules.push((r.prefix.clone(), up));
}
Ok(RouteTable::new(rules))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn repo_config_toml_parses_and_builds() {
let raw = include_str!("../../../config.toml");
let cfg: ConfigToml = toml::from_str(raw).expect("config.toml should parse");
build_table(&cfg).expect("build_table should succeed for config.toml");
}
#[test]
fn build_table_rejects_upstream_without_authority() {
let cfg = ConfigToml {
health_interval_secs: 5,
default_cooldown_secs: 30,
tenant_rps: 1000,
tls: TlsToml {
cert_path: "certs/server.crt".into(),
key_path: "certs/server.key".into(),
client_ca_path: "certs/ca-bundle.crt".into(),
},
jwt: JwtToml {
jwks_path: "certs/jwt-pub.pem".into(),
issuer: None,
audience: None,
},
routes: vec![RouteToml {
prefix: "/svc-a".into(),
upstream: "/no-authority-here".into(),
cooldown_secs: None,
}],
};
assert!(build_table(&cfg).is_err());
}
#[test]
fn build_table_rejects_zero_cooldown() {
let raw = include_str!("../../../config.toml").replace(
"upstream = \"http://localhost:8001\"",
"upstream = \"http://localhost:8001\"\ncooldown_secs = 0",
);
let cfg: ConfigToml = toml::from_str(&raw).unwrap();
assert!(build_table(&cfg).is_err());
}
fn base() -> String {
include_str!("../../../config.toml").to_string()
}
#[test]
fn shipped_config_has_default_limits() {
let (_, ext) = parse_config(&base()).unwrap();
let l = &ext.limits;
assert_eq!(l.max_request_body_bytes, 8 * 1024 * 1024);
assert_eq!(
(
l.request_body_timeout_secs,
l.upstream_timeout_secs,
l.tls_handshake_timeout_secs,
l.first_request_timeout_secs,
l.h2_max_concurrent_streams,
l.shutdown_drain_secs
),
(30, 30, 10, 10, 64, 25)
);
}
#[test]
fn limits_validation_names_the_key() {
for (line, key) in [
("max_request_body_bytes = 0", "max_request_body_bytes"),
(
"max_request_body_bytes = 1073741825",
"max_request_body_bytes",
),
("request_body_timeout_secs = 0", "request_body_timeout_secs"),
("upstream_timeout_secs = 86401", "upstream_timeout_secs"),
(
"tls_handshake_timeout_secs = 0",
"tls_handshake_timeout_secs",
),
(
"first_request_timeout_secs = 0",
"first_request_timeout_secs",
),
("h2_max_concurrent_streams = 0", "h2_max_concurrent_streams"),
("shutdown_drain_secs = 0", "shutdown_drain_secs"),
] {
let raw = format!("{}\n[limits]\n{line}\n", base());
let e = parse_config(&raw).unwrap_err().to_string();
assert!(e.contains(key), "{line}: {e}");
}
let ok = format!(
"{}\n[limits]\nmax_request_body_bytes = 1073741824\nupstream_timeout_secs = 86400\n",
base()
);
parse_config(&ok).unwrap();
}
#[test]
fn health_path_validation_and_mapping() {
for bad in ["healthz", "/h?x=1", "/h#f", "/a b", "/a\\u0001b"] {
let raw = base().replacen(
"upstream = \"http://localhost:8001\"",
&format!("upstream = \"http://localhost:8001\"\nhealth_path = \"{bad}\""),
1,
);
let e = parse_config(&raw).unwrap_err().to_string();
assert!(e.contains("health_path"), "{bad}: {e}");
}
let raw = base().replacen(
"upstream = \"http://localhost:8001\"",
"upstream = \"http://localhost:8001\"\nhealth_path = \"/ready\"\nhealth_disabled = true",
1,
);
let (cfg, ext) = parse_config(&raw).unwrap();
let t = build_table_ext(&cfg, &ext).unwrap();
let a = t.lookup("/svc-a").unwrap();
assert_eq!((a.health_path(), a.health_disabled()), ("/ready", true));
let b = t.lookup("/svc-b").unwrap();
assert_eq!((b.health_path(), b.health_disabled()), ("/health", false));
}
#[test]
fn ext_prefix_matching_no_route_is_an_error() {
let (cfg, _) = parse_config(&base()).unwrap();
let ext: ConfigExt =
toml::from_str("[[routes]]\nprefix = \"/typo\"\nhealth_disabled = true").unwrap();
let e = build_table_ext(&cfg, &ext).err().unwrap().to_string();
assert!(e.contains("/typo"), "{e}");
}
#[test]
fn zero_health_interval_rejected_everywhere() {
let raw = base().replace("health_interval_secs = 5", "health_interval_secs = 0");
assert!(parse_config(&raw).is_err());
let cfg: ConfigToml = toml::from_str(&raw).unwrap();
assert!(build_table(&cfg).is_err());
}
}