use serde::{Deserialize, Serialize};
use serde_json::Value;
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(deny_unknown_fields)]
pub struct ChannelConfig {
#[serde(default)]
pub rate_limit: Option<ChannelRateLimitConfig>,
#[serde(default)]
pub timeout_ms: Option<u64>,
#[serde(default)]
pub cache: Option<ChannelCacheConfig>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub origin_allow_list: Option<Vec<String>>,
#[serde(default)]
pub backpressure: Option<BackpressureConfig>,
#[serde(default)]
pub deduplication: Option<DeduplicationConfig>,
#[serde(default)]
pub validation_logic: Option<Value>,
#[serde(default)]
pub tracing: Option<ChannelTracingConfig>,
#[serde(default)]
pub response: Option<ChannelResponseConfig>,
#[serde(default)]
pub auth: Option<ChannelAuthConfig>,
}
impl ChannelConfig {
pub fn allowed_origins(&self) -> Option<&[String]> {
self.origin_allow_list.as_deref()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ChannelAuthConfig {
pub mode: AuthMode,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub keys: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub header: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub scheme: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub secret: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub signature_prefix: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum AuthMode {
ApiKey,
Hmac,
}
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
#[serde(default, deny_unknown_fields)]
pub struct ChannelResponseConfig {
pub mode: ResponseMode,
pub allowed_headers: Option<Vec<String>>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "lowercase")]
pub enum ResponseMode {
#[default]
Envelope,
Shaped,
}
pub const DEFAULT_ALLOWED_RESPONSE_HEADERS: &[&str] = &[
"content-type",
"location",
"cache-control",
"etag",
"last-modified",
"retry-after",
"content-language",
"link",
];
pub const FORBIDDEN_RESPONSE_HEADERS: &[&str] = &[
"connection",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"te",
"trailer",
"transfer-encoding",
"upgrade",
"content-length",
"x-request-id",
];
impl ChannelResponseConfig {
pub fn is_shaped(&self) -> bool {
self.mode == ResponseMode::Shaped
}
pub fn allows_header(&self, name: &str) -> bool {
if FORBIDDEN_RESPONSE_HEADERS.contains(&name) {
return false;
}
match self.allowed_headers {
Some(ref list) => list.iter().any(|h| h.eq_ignore_ascii_case(name)),
None => DEFAULT_ALLOWED_RESPONSE_HEADERS.contains(&name),
}
}
}
#[derive(Debug, Clone, Default, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ChannelTracingConfig {
#[serde(default)]
pub mode: Option<crate::config::TraceStorageMode>,
#[serde(default)]
pub sample_rate: Option<f64>,
#[serde(default)]
pub errors_only: Option<bool>,
#[serde(default)]
pub task_details: Option<bool>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, Default)]
#[serde(rename_all = "lowercase")]
pub enum BackendErrorPolicy {
#[default]
Allow,
Deny,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ChannelRateLimitConfig {
pub requests_per_second: u32,
#[serde(default)]
pub burst: Option<u32>,
#[serde(default)]
pub key_logic: Option<Value>,
#[serde(default)]
pub on_backend_error: BackendErrorPolicy,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct ChannelCacheConfig {
pub enabled: bool,
#[serde(default)]
pub ttl_secs: Option<u64>,
#[serde(default)]
pub cache_key_fields: Option<Vec<String>>,
#[serde(default)]
pub connector: Option<String>,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct BackpressureConfig {
pub max_concurrent_per_node: usize,
}
#[derive(Debug, Clone, Serialize, Deserialize)]
#[serde(deny_unknown_fields)]
pub struct DeduplicationConfig {
pub header: String,
#[serde(default)]
pub window_secs: Option<u64>,
#[serde(default)]
pub connector: Option<String>,
#[serde(default)]
pub on_backend_error: BackendErrorPolicy,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_channel_config_default() {
let config = ChannelConfig::default();
assert!(config.rate_limit.is_none());
assert!(config.timeout_ms.is_none());
assert!(config.cache.is_none());
assert!(config.backpressure.is_none());
assert!(config.deduplication.is_none());
assert!(config.validation_logic.is_none());
}
#[test]
fn test_channel_config_deserialization() {
let json = r#"{
"rate_limit": { "requests_per_second": 100, "burst": 20, "key_logic": { "var": "client_ip" } },
"timeout_ms": 5000,
"backpressure": { "max_concurrent_per_node": 200 },
"deduplication": { "header": "Idempotency-Key", "window_secs": 300 }
}"#;
let config: ChannelConfig = serde_json::from_str(json).expect("test");
let rl = config.rate_limit.expect("test");
assert_eq!(rl.requests_per_second, 100);
assert_eq!(rl.burst, Some(20));
assert!(rl.key_logic.is_some());
assert_eq!(rl.on_backend_error, BackendErrorPolicy::Allow);
assert_eq!(config.timeout_ms, Some(5000));
let bp = config.backpressure.expect("test");
assert_eq!(bp.max_concurrent_per_node, 200);
let dedup = config.deduplication.expect("test");
assert_eq!(dedup.header, "Idempotency-Key");
assert_eq!(dedup.window_secs, Some(300));
assert_eq!(dedup.on_backend_error, BackendErrorPolicy::Allow);
}
#[test]
fn test_backpressure_old_name_is_refused() {
let err =
serde_json::from_str::<ChannelConfig>(r#"{"backpressure": {"max_concurrent": 7}}"#)
.expect_err("the pre-1.0 `max_concurrent` spelling must not parse");
let message = err.to_string();
assert!(
message.contains("max_concurrent"),
"the error must name the offending key: {message}"
);
}
#[test]
fn test_on_backend_error_deserialization() {
let json = r#"{
"rate_limit": { "requests_per_second": 5, "on_backend_error": "deny" },
"deduplication": { "header": "idem", "on_backend_error": "deny" }
}"#;
let config: ChannelConfig = serde_json::from_str(json).expect("test");
assert_eq!(
config.rate_limit.expect("test").on_backend_error,
BackendErrorPolicy::Deny
);
assert_eq!(
config.deduplication.expect("test").on_backend_error,
BackendErrorPolicy::Deny
);
assert!(
serde_json::from_str::<ChannelConfig>(
r#"{"deduplication": {"header": "idem", "on_backend_error": "explode"}}"#
)
.is_err(),
"unknown policy values must be rejected, not defaulted"
);
}
#[test]
fn test_origin_allow_list() {
let new_key: ChannelConfig =
serde_json::from_str(r#"{"origin_allow_list": ["https://app.example.com"]}"#)
.expect("test");
assert_eq!(
new_key.allowed_origins(),
Some(["https://app.example.com".to_string()].as_slice())
);
assert!(ChannelConfig::default().allowed_origins().is_none());
}
#[test]
fn test_pre_1_0_cors_spelling_is_refused_not_ignored() {
for stored in [
r#"{"cors": {"allowed_origins": ["https://old.example.com"]}}"#,
r#"{"cors": {}}"#,
] {
let err = serde_json::from_str::<ChannelConfig>(stored)
.expect_err("the pre-1.0 `cors` spelling must not parse");
let message = err.to_string();
assert!(
message.contains("cors"),
"the error must name the offending key: {message}"
);
}
}
#[test]
fn test_unknown_channel_config_key_is_refused() {
let err = serde_json::from_str::<ChannelConfig>(
r#"{"deduplicaton": {"header": "Idempotency-Key"}}"#,
)
.expect_err("a misspelled guard key must not be silently ignored");
assert!(
err.to_string().contains("deduplicaton"),
"the error must name the typo: {err}"
);
}
#[test]
fn test_unknown_key_inside_a_guard_is_refused() {
for (config, typo) in [
(
r#"{"rate_limit": {"requests_per_second": 10, "key_logic_": {"var": "client_ip"}}}"#,
"key_logic_",
),
(
r#"{"deduplication": {"header": "Idempotency-Key", "window_seconds": 60}}"#,
"window_seconds",
),
(
r#"{"cache": {"enabled": true, "ttl_seconds": 30}}"#,
"ttl_seconds",
),
(r#"{"tracing": {"sampling_rate": 0.5}}"#, "sampling_rate"),
] {
let err = serde_json::from_str::<ChannelConfig>(config)
.expect_err("a misspelled key inside a guard must not be silently ignored");
assert!(
err.to_string().contains(typo),
"the error must name the typo `{typo}`: {err}"
);
}
}
#[test]
fn test_channel_config_empty_json() {
let config: ChannelConfig = serde_json::from_str("{}").expect("test");
assert!(config.rate_limit.is_none());
assert!(config.timeout_ms.is_none());
}
}