Skip to main content

praxis_protocol/tcp/
service.rs

1// SPDX-License-Identifier: Apache-2.0
2// Copyright (c) 2024 Praxis Contributors
3
4//! TCP protocol service construction and registration.
5
6use std::{sync::Arc, time::Duration};
7
8use arc_swap::ArcSwap;
9use pingora_core::services::listening::Service;
10use praxis_core::{ProxyError, config::Config};
11use praxis_filter::{FilterPipeline, FilterRegistry};
12use tokio::sync::{Semaphore, watch};
13
14use super::{proxy, tls};
15use crate::{ListenerPipelines, Protocol};
16
17// -----------------------------------------------------------------------------
18// PingoraTcp
19// -----------------------------------------------------------------------------
20
21/// Pingora-backed raw TCP/L4 protocol implementation.
22///
23/// Groups TCP listeners by `(upstream address, idle timeout, max duration)`,
24/// creating one bidirectional forwarder per unique combination. Implements [`Protocol`].
25///
26/// [`Protocol`]: crate::Protocol
27pub struct PingoraTcp;
28
29impl Protocol for PingoraTcp {
30    fn register(
31        self: Box<Self>,
32        server: &mut praxis_core::PingoraServerRuntime,
33        config: &Config,
34        pipelines: &ListenerPipelines,
35    ) -> Result<Vec<watch::Sender<bool>>, ProxyError> {
36        let groups = tls::group_tcp_listeners(config);
37        tls::validate_tcp_group_consistency(&groups)?;
38        #[expect(clippy::expect_used, reason = "empty pipeline is infallible")]
39        let fallback_pipeline = Arc::new(ArcSwap::from_pointee(
40            FilterPipeline::build(&mut [], &FilterRegistry::with_builtins()).expect("empty pipeline is valid"),
41        ));
42
43        let mut cert_watcher_shutdowns = Vec::new();
44        for (group_key, listeners) in groups {
45            let mut service = build_tcp_service(&group_key, &listeners, pipelines, &fallback_pipeline, config);
46            cert_watcher_shutdowns.extend(tls::register_tcp_listeners(
47                &mut service,
48                &listeners,
49                group_key.0.as_deref(),
50            )?);
51            server.server_mut().add_service(service);
52        }
53
54        Ok(cert_watcher_shutdowns)
55    }
56}
57
58// -----------------------------------------------------------------------------
59// Service Construction
60// -----------------------------------------------------------------------------
61
62/// The listener group key: shared upstream address, cluster, idle timeout
63/// (ms), and max session duration (secs).
64type TcpGroupKey = (Option<String>, Option<String>, Option<u64>, Option<u64>);
65
66/// Build the TCP proxy service for one listener group.
67fn build_tcp_service(
68    group_key: &TcpGroupKey,
69    listeners: &[&praxis_core::config::Listener],
70    pipelines: &ListenerPipelines,
71    fallback_pipeline: &Arc<ArcSwap<FilterPipeline>>,
72    config: &Config,
73) -> Service<proxy::PingoraTcpProxy> {
74    let (upstream_opt, cluster_opt, timeout_ms, max_dur_secs) = group_key;
75    let pipeline = listeners
76        .first()
77        .and_then(|l| pipelines.get(&l.name))
78        .map_or_else(|| Arc::clone(fallback_pipeline), Arc::clone);
79    let session_timeout = timeout_ms.map(Duration::from_millis);
80    let max_duration = max_dur_secs.map(Duration::from_secs);
81    let connection_semaphore = listeners
82        .first()
83        .and_then(|l| l.max_connections)
84        .map(|max| Arc::new(Semaphore::new(max as usize)));
85    let (listener_names, default_listener_name) = build_listener_labels(listeners);
86    let app = proxy::PingoraTcpProxy::new(
87        upstream_opt.clone(),
88        cluster_opt.clone().map(Arc::from),
89        pipeline,
90        session_timeout,
91        max_duration,
92        connection_semaphore,
93        config.insecure_options.allow_private_upstreams,
94        listener_names,
95        default_listener_name,
96    );
97    Service::new(tcp_service_name(upstream_opt.as_deref(), cluster_opt.as_deref()), app)
98}
99
100/// Derive the Pingora service name for a TCP listener group.
101fn tcp_service_name(upstream: Option<&str>, cluster: Option<&str>) -> String {
102    match (upstream, cluster) {
103        (Some(addr), _) => format!("tcp-proxy:{addr}"),
104        (_, Some(cluster)) => format!("tcp-proxy:cluster:{cluster}"),
105        _ => "tcp-proxy:filter-routed".to_owned(),
106    }
107}
108
109/// Build the per-address metric label map and the default listener label.
110fn build_listener_labels(
111    listeners: &[&praxis_core::config::Listener],
112) -> (
113    std::collections::HashMap<String, ::metrics::SharedString>,
114    ::metrics::SharedString,
115) {
116    // `from_shared` keeps labels as refcounted `Arc<str>`s: they are cloned
117    // per connection, and owned `String` labels would deep-copy on every clone.
118    let by_address = listeners
119        .iter()
120        .map(|l| {
121            (
122                l.address.clone(),
123                ::metrics::SharedString::from_shared(Arc::from(l.name.as_str())),
124            )
125        })
126        .collect();
127    let default = listeners.first().map_or_else(
128        || ::metrics::SharedString::const_str("unknown"),
129        |l| ::metrics::SharedString::from_shared(Arc::from(l.name.as_str())),
130    );
131    (by_address, default)
132}