use std::collections::HashMap;
use pingora_core::services::listening::Service;
use praxis_core::{
ProxyError,
config::{Config, ProtocolKind},
};
use tracing::info;
use super::proxy::PingoraTcpProxy;
pub(super) type TcpGroupKey = (Option<String>, Option<String>, Option<u64>, Option<u64>);
pub(super) fn group_tcp_listeners(config: &Config) -> HashMap<TcpGroupKey, Vec<&praxis_core::config::Listener>> {
let mut groups: HashMap<TcpGroupKey, Vec<&praxis_core::config::Listener>> = HashMap::new();
for listener in &config.listeners {
if listener.protocol != ProtocolKind::Tcp {
continue;
}
let key = (
listener.upstream.clone(),
listener.cluster.clone(),
listener.tcp_session_timeout_ms,
listener.tcp_max_duration_secs,
);
groups.entry(key).or_default().push(listener);
}
groups
}
pub(super) fn validate_tcp_group_consistency(
groups: &HashMap<TcpGroupKey, Vec<&praxis_core::config::Listener>>,
) -> Result<(), ProxyError> {
for listeners in groups.values() {
let Some((first, rest)) = listeners.split_first() else {
continue;
};
for listener in rest {
if listener.filter_chains != first.filter_chains {
return Err(ProxyError::Config(format!(
"TCP listeners '{}' and '{}' share the same upstream/timeout \
group but have different filter_chains; grouped listeners \
must use identical chains",
first.name, listener.name
)));
}
if listener.max_connections != first.max_connections {
return Err(ProxyError::Config(format!(
"TCP listeners '{}' and '{}' share the same upstream/timeout \
group but have different max_connections; grouped listeners \
must use identical limits",
first.name, listener.name
)));
}
}
}
Ok(())
}
pub fn validate_tcp_groups(config: &Config) -> Result<(), ProxyError> {
validate_tcp_group_consistency(&group_tcp_listeners(config))
}
pub(super) fn register_tcp_listeners(
service: &mut Service<PingoraTcpProxy>,
listeners: &[&praxis_core::config::Listener],
upstream: Option<&str>,
) -> Result<Vec<tokio::sync::watch::Sender<bool>>, ProxyError> {
let display_upstream = upstream.unwrap_or("filter-routed");
let mut shutdown_senders = Vec::new();
for listener in listeners {
if let Some(tls) = &listener.tls {
let (tls_settings, watcher_shutdown) = build_tcp_tls_settings(tls, &listener.address)?;
if let Some(tx) = watcher_shutdown {
shutdown_senders.push(tx);
}
service.add_tls_with_settings(&listener.address, None, tls_settings);
} else {
service.add_tcp(&listener.address);
}
info!(
name = %listener.name,
address = %listener.address,
upstream = %display_upstream,
"TCP listener registered"
);
}
Ok(shutdown_senders)
}
fn build_tcp_tls_settings(
tls: &praxis_tls::ListenerTls,
address: &str,
) -> Result<
(
pingora_core::listeners::tls::TlsSettings,
Option<tokio::sync::watch::Sender<bool>>,
),
ProxyError,
> {
crate::tls_setup::build_tls_settings(tls, address, "TCP", false)
}
#[cfg(test)]
#[expect(clippy::allow_attributes, reason = "blanket test suppressions")]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::indexing_slicing, reason = "tests")]
mod tests {
use praxis_core::config::{AdminConfig, BodyLimitsConfig, InsecureOptions, RuntimeConfig};
use super::*;
#[test]
fn group_tcp_listeners_groups_by_upstream_and_timeout() {
let config = Config::from_yaml(
r#"
listeners:
- name: db1
address: "0.0.0.0:5432"
protocol: tcp
upstream: "10.0.0.1:5432"
- name: db2
address: "0.0.0.0:5433"
protocol: tcp
upstream: "10.0.0.1:5432"
"#,
)
.unwrap();
let groups = group_tcp_listeners(&config);
assert_eq!(groups.len(), 1, "same upstream + timeout should produce one group");
let default_timeout = config.listeners[0].tcp_session_timeout_ms;
let key = (Some("10.0.0.1:5432".to_owned()), None, default_timeout, None);
assert_eq!(groups[&key].len(), 2, "both listeners should be in the same group");
}
#[test]
fn validate_tcp_groups_rejects_divergent_filter_chains() {
let config = Config::from_yaml(
r#"
listeners:
- name: db1
address: "0.0.0.0:5432"
protocol: tcp
upstream: "10.0.0.1:5432"
filter_chains: [a]
- name: db2
address: "0.0.0.0:5433"
protocol: tcp
upstream: "10.0.0.1:5432"
filter_chains: [b]
filter_chains:
- name: a
filters: []
- name: b
filters: []
"#,
)
.unwrap();
let err = validate_tcp_groups(&config).unwrap_err().to_string();
assert!(
err.contains("db1") && err.contains("db2") && err.contains("filter_chains"),
"error should name both listeners and the field: {err}"
);
}
#[test]
fn validate_tcp_groups_accepts_consistent_groups() {
let config = Config::from_yaml(
r#"
listeners:
- name: db1
address: "0.0.0.0:5432"
protocol: tcp
upstream: "10.0.0.1:5432"
- name: db2
address: "0.0.0.0:5433"
protocol: tcp
upstream: "10.0.0.1:5432"
"#,
)
.unwrap();
assert!(
validate_tcp_groups(&config).is_ok(),
"identical grouped listeners should validate"
);
}
#[test]
fn group_tcp_listeners_separates_different_upstreams() {
let config = Config::from_yaml(
r#"
listeners:
- name: db
address: "0.0.0.0:5432"
protocol: tcp
upstream: "10.0.0.1:5432"
- name: cache
address: "0.0.0.0:6379"
protocol: tcp
upstream: "10.0.0.2:6379"
"#,
)
.unwrap();
let groups = group_tcp_listeners(&config);
assert_eq!(groups.len(), 2, "different upstreams should produce separate groups");
}
#[test]
fn group_tcp_listeners_separates_different_timeouts() {
let config = Config::from_yaml(
r#"
listeners:
- name: a
address: "0.0.0.0:5432"
protocol: tcp
upstream: "10.0.0.1:5432"
- name: b
address: "0.0.0.0:5433"
protocol: tcp
upstream: "10.0.0.1:5432"
tcp_session_timeout_ms: 30000
"#,
)
.unwrap();
let groups = group_tcp_listeners(&config);
assert_eq!(
groups.len(),
2,
"same upstream but different timeouts should produce separate groups"
);
}
#[test]
fn group_tcp_listeners_skips_http_listeners() {
let config = config_with_http_and_tcp();
let groups = group_tcp_listeners(&config);
assert_eq!(groups.len(), 1, "HTTP listeners should be excluded");
let timeout = config
.listeners
.iter()
.find(|l| l.protocol == ProtocolKind::Tcp)
.unwrap()
.tcp_session_timeout_ms;
let key = (Some("10.0.0.1:5432".to_owned()), None, timeout, None);
assert!(groups.contains_key(&key), "only TCP listener should be grouped");
}
#[test]
fn group_tcp_listeners_includes_tcp_without_upstream() {
let config = config_with_tcp_no_upstream();
let groups = group_tcp_listeners(&config);
assert_eq!(
groups.len(),
1,
"TCP listener without upstream should be grouped with None key"
);
let key = (None, None, None, None);
assert!(groups.contains_key(&key), "group key should have None upstream");
}
#[test]
fn group_tcp_listeners_http_only_yields_empty() {
let config = config_http_only();
let groups = group_tcp_listeners(&config);
assert!(
groups.is_empty(),
"config with only HTTP listeners should yield empty groups"
);
}
#[test]
fn validate_consistency_accepts_matching_chains() {
let config = Config::from_yaml(
r#"
listeners:
- name: a
address: "0.0.0.0:5432"
protocol: tcp
upstream: "10.0.0.1:5432"
- name: b
address: "0.0.0.0:5433"
protocol: tcp
upstream: "10.0.0.1:5432"
"#,
)
.unwrap();
let groups = group_tcp_listeners(&config);
assert!(
validate_tcp_group_consistency(&groups).is_ok(),
"identical chains should pass consistency check"
);
}
#[test]
fn validate_consistency_rejects_different_chains() {
let owned = groups_with_different_filter_chains();
let groups = to_ref_groups(&owned);
let err = validate_tcp_group_consistency(&groups).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("different filter_chains"),
"error should mention filter_chains: {msg}"
);
}
#[test]
fn validate_consistency_rejects_different_max_connections() {
let owned = groups_with_different_max_connections();
let groups = to_ref_groups(&owned);
let err = validate_tcp_group_consistency(&groups).unwrap_err();
let msg = err.to_string();
assert!(
msg.contains("different max_connections"),
"error should mention max_connections: {msg}"
);
}
#[test]
fn validate_consistency_accepts_single_listener_groups() {
let config = Config::from_yaml(
r#"
listeners:
- name: solo
address: "0.0.0.0:5432"
protocol: tcp
upstream: "10.0.0.1:5432"
"#,
)
.unwrap();
let groups = group_tcp_listeners(&config);
assert!(
validate_tcp_group_consistency(&groups).is_ok(),
"single-listener group should always pass"
);
}
#[test]
fn validate_consistency_accepts_empty_groups() {
let groups: HashMap<TcpGroupKey, Vec<&praxis_core::config::Listener>> = HashMap::new();
assert!(
validate_tcp_group_consistency(&groups).is_ok(),
"empty groups should pass"
);
}
fn config_with_http_and_tcp() -> Config {
Config::from_yaml(
r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
- name: db
address: "0.0.0.0:5432"
protocol: tcp
upstream: "10.0.0.1:5432"
filter_chains:
- name: main
filters:
- filter: router
routes:
- path_prefix: "/"
cluster: default
- filter: load_balancer
clusters:
- name: default
endpoints: ["127.0.0.1:9090"]
insecure_options:
allow_private_endpoints: true
"#,
)
.unwrap()
}
fn config_http_only() -> Config {
Config::from_yaml(
r#"
listeners:
- name: web
address: "0.0.0.0:8080"
filter_chains: [main]
filter_chains:
- name: main
filters:
- filter: router
routes:
- path_prefix: "/"
cluster: default
- filter: load_balancer
clusters:
- name: default
endpoints: ["127.0.0.1:9090"]
insecure_options:
allow_private_endpoints: true
"#,
)
.unwrap()
}
fn config_with_tcp_no_upstream() -> Config {
use praxis_core::config::{Listener, MetricsConfig, TelemetryConfig};
Config {
admin: AdminConfig::default(),
body_limits: BodyLimitsConfig::default(),
clusters: vec![],
filter_chains: vec![],
insecure_options: InsecureOptions::default(),
listeners: vec![Listener {
name: "orphan".to_owned(),
address: "0.0.0.0:9999".to_owned(),
cluster: None,
downstream_keepalive_timeout_ms: None,
downstream_read_timeout_ms: None,
filter_chains: vec![],
max_connections: None,
protocol: ProtocolKind::Tcp,
tcp_session_timeout_ms: None,
tcp_max_duration_secs: None,
tls: None,
upstream: None,
}],
metrics: MetricsConfig::default(),
runtime: RuntimeConfig::default(),
shutdown_timeout_secs: 10,
telemetry: TelemetryConfig::default(),
}
}
fn to_ref_groups(
owned: &HashMap<TcpGroupKey, Vec<praxis_core::config::Listener>>,
) -> HashMap<TcpGroupKey, Vec<&praxis_core::config::Listener>> {
owned.iter().map(|(k, v)| (k.clone(), v.iter().collect())).collect()
}
fn groups_with_different_filter_chains() -> HashMap<TcpGroupKey, Vec<praxis_core::config::Listener>> {
let mut a = make_tcp_listener("a", "0.0.0.0:5432");
a.filter_chains = vec!["chain-a".to_owned()];
let mut b = make_tcp_listener("b", "0.0.0.0:5433");
b.filter_chains = vec!["chain-b".to_owned()];
make_group(vec![a, b])
}
fn groups_with_different_max_connections() -> HashMap<TcpGroupKey, Vec<praxis_core::config::Listener>> {
let mut a = make_tcp_listener("a", "0.0.0.0:5432");
a.max_connections = Some(100);
let mut b = make_tcp_listener("b", "0.0.0.0:5433");
b.max_connections = Some(200);
make_group(vec![a, b])
}
fn make_tcp_listener(name: &str, address: &str) -> praxis_core::config::Listener {
use praxis_core::config::Listener;
Listener {
name: name.to_owned(),
address: address.to_owned(),
cluster: None,
downstream_keepalive_timeout_ms: None,
downstream_read_timeout_ms: None,
filter_chains: vec![],
max_connections: None,
protocol: ProtocolKind::Tcp,
tcp_session_timeout_ms: None,
tcp_max_duration_secs: None,
tls: None,
upstream: Some("10.0.0.1:5432".to_owned()),
}
}
fn make_group(
listeners: Vec<praxis_core::config::Listener>,
) -> HashMap<TcpGroupKey, Vec<praxis_core::config::Listener>> {
let key = (Some("10.0.0.1:5432".to_owned()), None, None, None);
let mut groups = HashMap::new();
groups.insert(key, listeners);
groups
}
}