use serde::de::{self, Visitor};
use serde::{Deserialize, Deserializer, Serialize, Serializer};
use std::collections::BTreeMap;
use std::fmt;
use std::time::Duration;
fn default_require_auth_by_default() -> bool {
true
}
fn default_throttle_status() -> u16 {
429
}
fn default_body_limit_bytes() -> usize {
16 * 1024 * 1024
}
fn default_healthcheck_timeout_ms() -> u64 {
500
}
fn default_gateway_sync_interval_secs() -> u64 {
10
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
#[allow(clippy::struct_excessive_bools)]
pub struct ApiGatewayConfig {
pub bind_addr: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub advertise_uri: Option<String>,
#[serde(default)]
pub enable_docs: bool,
#[serde(default)]
pub cors_enabled: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub cors: Option<CorsConfig>,
#[serde(default)]
pub openapi: OpenApiConfig,
#[serde(default)]
pub defaults: Defaults,
#[serde(default)]
pub auth_disabled: bool,
#[serde(default = "default_require_auth_by_default")]
pub require_auth_by_default: bool,
#[serde(default)]
pub prefix_path: String,
#[serde(default)]
pub route_policies: RoutePoliciesConfig,
#[serde(default)]
pub metrics: MetricsConfig,
#[serde(default = "default_healthcheck_timeout_ms")]
pub healthcheck_timeout_ms: u64,
#[serde(default)]
pub health: HealthConfig,
#[serde(default)]
pub gateway_proxy: GatewayProxyConfig,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub internal_auth: Option<toolkit_security::InternalAuthConfig>,
#[serde(default)]
pub rate_limit_zones: BTreeMap<String, RateLimitZone>,
#[serde(default)]
pub in_flight_limit_zones: BTreeMap<String, InFlightLimitZone>,
#[serde(default)]
pub trusted_proxy_hops: usize,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields, default)]
pub struct GatewayProxyConfig {
pub enabled: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub directory_endpoint: Option<String>,
pub sync_interval_secs: u64,
}
impl Default for GatewayProxyConfig {
fn default() -> Self {
Self {
enabled: false,
directory_endpoint: None,
sync_interval_secs: default_gateway_sync_interval_secs(),
}
}
}
impl GatewayProxyConfig {
pub fn validate(&self) -> Result<(), String> {
if !self.enabled {
return Ok(());
}
match self.directory_endpoint.as_deref() {
Some(endpoint) if !endpoint.trim().is_empty() => Ok(()),
_ => Err(
"invalid gateway_proxy configuration: `gateway_proxy.enabled` is true but \
`gateway_proxy.directory_endpoint` is unset or empty; set it to the \
DirectoryService gRPC endpoint (e.g. \"http://gear-orchestrator:50051\") \
or disable the reverse proxy"
.to_owned(),
),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize, Default)]
#[serde(rename_all = "lowercase")]
pub enum HealthServeMode {
#[default]
Main,
Separate,
Both,
}
#[derive(Debug, Clone, Deserialize, Serialize, Default)]
#[serde(deny_unknown_fields, default)]
pub struct HealthConfig {
pub serve: HealthServeMode,
#[serde(skip_serializing_if = "Option::is_none")]
pub bind_addr: Option<String>,
}
impl Default for ApiGatewayConfig {
fn default() -> Self {
Self {
bind_addr: String::default(),
advertise_uri: None,
enable_docs: false,
cors_enabled: false,
cors: None,
openapi: OpenApiConfig::default(),
defaults: Defaults::default(),
auth_disabled: false,
require_auth_by_default: default_require_auth_by_default(),
prefix_path: String::default(),
route_policies: RoutePoliciesConfig::default(),
metrics: MetricsConfig::default(),
healthcheck_timeout_ms: default_healthcheck_timeout_ms(),
health: HealthConfig::default(),
gateway_proxy: GatewayProxyConfig::default(),
internal_auth: None,
rate_limit_zones: BTreeMap::new(),
in_flight_limit_zones: BTreeMap::new(),
trusted_proxy_hops: 0,
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields, default)]
pub struct Defaults {
pub body_limit_bytes: usize,
}
impl Default for Defaults {
fn default() -> Self {
Self {
body_limit_bytes: default_body_limit_bytes(),
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields, default)]
pub struct CorsConfig {
pub allowed_origins: Vec<String>,
pub allowed_methods: Vec<String>,
pub allowed_headers: Vec<String>,
pub exposed_headers: Vec<String>,
pub allow_credentials: bool,
pub max_age_seconds: u64,
}
impl Default for CorsConfig {
fn default() -> Self {
Self {
allowed_origins: vec!["*".to_owned()],
allowed_methods: vec![
"GET".to_owned(),
"POST".to_owned(),
"PUT".to_owned(),
"PATCH".to_owned(),
"DELETE".to_owned(),
"OPTIONS".to_owned(),
],
allowed_headers: vec!["*".to_owned()],
exposed_headers: vec!["ETag".to_owned()],
allow_credentials: false,
max_age_seconds: 600,
}
}
}
#[derive(Debug, Clone, Deserialize, Serialize, Default)]
#[serde(deny_unknown_fields, default)]
pub struct MetricsConfig {
pub prefix: String,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields, default)]
pub struct OpenApiConfig {
pub title: String,
pub version: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
}
impl Default for OpenApiConfig {
fn default() -> Self {
Self {
title: "API Documentation".to_owned(),
version: "0.1.0".to_owned(),
description: None,
}
}
}
#[derive(Debug, Clone, Default, Deserialize, Serialize)]
#[serde(deny_unknown_fields, default)]
pub struct RoutePoliciesConfig {
pub enabled: bool,
pub rules: Vec<RoutePolicyRule>,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct RoutePolicyRule {
pub path: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub method: Option<String>,
pub required_scopes: Vec<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RateSpec {
pub rps: u32,
}
impl<'de> Deserialize<'de> for RateSpec {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let raw = String::deserialize(deserializer)?;
let s = raw.trim();
let value = s.strip_suffix("/s").ok_or_else(|| {
de::Error::custom(format!(
"invalid rate '{s}': only the '/s' (per-second) unit is supported, e.g. '50/s'"
))
})?;
let rps: u32 = value.trim().parse().map_err(|_| {
de::Error::custom(format!(
"invalid rate '{s}': '{value}' is not a valid integer"
))
})?;
Ok(Self { rps })
}
}
impl Serialize for RateSpec {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
serializer.serialize_str(&format!("{}/s", self.rps))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum RetryAfter {
#[default]
Auto,
Seconds(u64),
}
impl<'de> Deserialize<'de> for RetryAfter {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
struct RetryAfterVisitor;
impl Visitor<'_> for RetryAfterVisitor {
type Value = RetryAfter;
fn expecting(&self, f: &mut fmt::Formatter) -> fmt::Result {
f.write_str("the string \"auto\" or a non-negative integer number of seconds")
}
fn visit_u64<E: de::Error>(self, v: u64) -> Result<Self::Value, E> {
Ok(RetryAfter::Seconds(v))
}
fn visit_i64<E: de::Error>(self, v: i64) -> Result<Self::Value, E> {
u64::try_from(v)
.map(RetryAfter::Seconds)
.map_err(|_| E::custom("response_retry_after must be non-negative"))
}
fn visit_str<E: de::Error>(self, v: &str) -> Result<Self::Value, E> {
if v.eq_ignore_ascii_case("auto") {
Ok(RetryAfter::Auto)
} else {
Err(E::custom(format!(
"invalid response_retry_after '{v}': expected \"auto\" or a number of seconds"
)))
}
}
}
deserializer.deserialize_any(RetryAfterVisitor)
}
}
impl Serialize for RetryAfter {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
match self {
Self::Auto => serializer.serialize_str("auto"),
Self::Seconds(n) => serializer.serialize_u64(*n),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Deserialize, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum KeyType {
Identity,
Ip,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct KeyConfig {
#[serde(rename = "type")]
pub key_type: KeyType,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct RateLimitZone {
pub rate_limit: RateSpec,
pub burst_limit: u32,
#[serde(default = "default_throttle_status")]
pub response_status_code: u16,
#[serde(default)]
pub response_retry_after: RetryAfter,
pub key: KeyConfig,
pub max_keys: u64,
}
#[derive(Debug, Clone, Deserialize, Serialize)]
#[serde(deny_unknown_fields)]
pub struct InFlightLimitZone {
pub in_flight_limit: u32,
pub backlog_limit: u32,
#[serde(with = "humantime_serde")]
pub backlog_timeout: Duration,
#[serde(default = "default_throttle_status")]
pub response_status_code: u16,
pub key: KeyConfig,
pub max_keys: u64,
#[serde(default)]
pub excluded_keys: Vec<String>,
}
impl ApiGatewayConfig {
pub fn validate_throttling(&self) -> anyhow::Result<()> {
for (name, zone) in &self.rate_limit_zones {
if zone.rate_limit.rps == 0 {
anyhow::bail!("rate_limit_zone '{name}': rate_limit must be greater than 0");
}
if zone.burst_limit == 0 {
anyhow::bail!("rate_limit_zone '{name}': burst_limit must be greater than 0");
}
if zone.max_keys == 0 {
anyhow::bail!("rate_limit_zone '{name}': max_keys must be greater than 0");
}
validate_status(name, zone.response_status_code)?;
}
for (name, zone) in &self.in_flight_limit_zones {
if zone.in_flight_limit == 0 {
anyhow::bail!(
"in_flight_limit_zone '{name}': in_flight_limit must be greater than 0"
);
}
if zone.max_keys == 0 {
anyhow::bail!("in_flight_limit_zone '{name}': max_keys must be greater than 0");
}
validate_status(name, zone.response_status_code)?;
}
Ok(())
}
}
fn validate_status(zone: &str, status: u16) -> anyhow::Result<()> {
let code = http::StatusCode::from_u16(status)
.map_err(|_| anyhow::anyhow!("zone '{zone}': invalid response_status_code {status}"))?;
if !(code.is_client_error() || code.is_server_error()) {
anyhow::bail!(
"zone '{zone}': response_status_code {status} must be a 4xx or 5xx error status"
);
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::GatewayProxyConfig;
#[test]
fn disabled_proxy_needs_no_endpoint() {
let cfg = GatewayProxyConfig::default();
assert!(!cfg.enabled);
assert!(cfg.validate().is_ok());
}
#[test]
fn enabled_proxy_without_endpoint_is_rejected() {
let cfg = GatewayProxyConfig {
enabled: true,
directory_endpoint: None,
..GatewayProxyConfig::default()
};
assert!(cfg.validate().is_err());
}
#[test]
fn enabled_proxy_with_blank_endpoint_is_rejected() {
let cfg = GatewayProxyConfig {
enabled: true,
directory_endpoint: Some(" ".to_owned()),
..GatewayProxyConfig::default()
};
assert!(cfg.validate().is_err());
}
#[test]
fn enabled_proxy_with_endpoint_is_accepted() {
let cfg = GatewayProxyConfig {
enabled: true,
directory_endpoint: Some("http://gear-orchestrator:50051".to_owned()),
..GatewayProxyConfig::default()
};
assert!(cfg.validate().is_ok());
}
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
mod throttling_tests {
use super::*;
const FULL_CONFIG: &str = r#"
bind_addr: "0.0.0.0:8086"
rate_limit_zones:
rl_identity:
rate_limit: 50/s
burst_limit: 100
response_status_code: 429
response_retry_after: auto
key:
type: identity
max_keys: 50000
in_flight_limit_zones:
ifl_identity:
in_flight_limit: 64
backlog_limit: 128
backlog_timeout: 30s
response_status_code: 429
key:
type: identity
max_keys: 50000
ifl_identity_expensive_op:
in_flight_limit: 4
backlog_limit: 8
backlog_timeout: 30s
response_status_code: 429
key:
type: identity
max_keys: 50000
excluded_keys:
- "150853ab-322c-455d-9793-8d71bf6973d9"
"#;
fn parse(yaml: &str) -> ApiGatewayConfig {
serde_saphyr::from_str(yaml).expect("config should deserialize")
}
#[test]
fn parses_rate_limit_zones() {
let cfg = parse(FULL_CONFIG);
assert_eq!(cfg.rate_limit_zones.len(), 1);
let rl = &cfg.rate_limit_zones["rl_identity"];
assert_eq!(rl.rate_limit, RateSpec { rps: 50 });
assert_eq!(rl.burst_limit, 100);
assert_eq!(rl.response_status_code, 429);
assert_eq!(rl.response_retry_after, RetryAfter::Auto);
assert_eq!(rl.key.key_type, KeyType::Identity);
assert_eq!(rl.max_keys, 50000);
}
#[test]
fn parses_in_flight_zones() {
let cfg = parse(FULL_CONFIG);
assert_eq!(cfg.in_flight_limit_zones.len(), 2);
let ifl = &cfg.in_flight_limit_zones["ifl_identity"];
assert_eq!(ifl.in_flight_limit, 64);
assert_eq!(ifl.backlog_limit, 128);
assert_eq!(ifl.backlog_timeout, Duration::from_secs(30));
assert!(ifl.excluded_keys.is_empty());
let expensive = &cfg.in_flight_limit_zones["ifl_identity_expensive_op"];
assert_eq!(expensive.in_flight_limit, 4);
assert_eq!(
expensive.excluded_keys,
vec!["150853ab-322c-455d-9793-8d71bf6973d9".to_owned()]
);
}
#[test]
fn full_config_validates() {
parse(FULL_CONFIG)
.validate_throttling()
.expect("config is valid");
}
#[test]
fn rate_spec_rejects_non_per_second_unit() {
let err = serde_saphyr::from_str::<RateSpec>("50/m")
.unwrap_err()
.to_string();
assert!(err.contains("/s"), "unexpected error: {err}");
}
#[test]
fn rate_spec_rejects_non_integer() {
assert!(serde_saphyr::from_str::<RateSpec>("abc/s").is_err());
}
#[test]
fn retry_after_parses_auto_and_seconds() {
assert_eq!(
serde_saphyr::from_str::<RetryAfter>("auto").unwrap(),
RetryAfter::Auto
);
assert_eq!(
serde_saphyr::from_str::<RetryAfter>("15").unwrap(),
RetryAfter::Seconds(15)
);
}
#[test]
fn validate_status_accepts_4xx_5xx_rejects_others() {
assert!(validate_status("z", 429).is_ok());
assert!(validate_status("z", 503).is_ok());
let err = validate_status("zone_a", 200).unwrap_err().to_string();
assert!(err.contains("zone_a") && err.contains("200"), "got: {err}");
assert!(validate_status("z", 302).is_err());
assert!(validate_status("z", 999).is_err());
}
#[test]
fn validation_rejects_zero_limit() {
let yaml = r#"
bind_addr: "0.0.0.0:8086"
rate_limit_zones:
bad:
rate_limit: 0/s
burst_limit: 100
key:
type: identity
max_keys: 100
"#;
let cfg = parse(yaml);
assert!(cfg.validate_throttling().is_err());
}
#[test]
fn empty_config_validates() {
let cfg = ApiGatewayConfig::default();
cfg.validate_throttling().expect("empty config is valid");
}
}