Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
123 changes: 114 additions & 9 deletions gateway/src/core/net.rs
Original file line number Diff line number Diff line change
Expand Up @@ -102,6 +102,29 @@ fn loopback_flag_enabled(value: Option<&str>) -> Result<bool, String> {
}
}

/// A process-wide pooled HTTP client for vendor calls, one per `key`, built by `build` on first
/// use and reused after.
///
/// A `reqwest::Client` IS its connection pool, so building one per request -- which the audio
/// paths did -- throws the pool away every time: each call paid DNS, TCP and TLS to the vendor.
/// Reusing it keeps keep-alive connections to each vendor host (one pool per host inside the
/// client), which is per-deployment reuse without keying on deployments. Key by whatever changes
/// the client's configuration (schemes, timeout). Build these HTTP/1.1-only: many long requests
/// on one shared HTTP/2 connection queue behind the vendor's concurrent-stream limit.
pub fn shared_http_client(
key: &str,
build: impl FnOnce() -> Result<reqwest::Client, reqwest::Error>,
) -> Result<reqwest::Client, reqwest::Error> {
static CLIENTS: std::sync::LazyLock<dashmap::DashMap<String, reqwest::Client>> =
std::sync::LazyLock::new(dashmap::DashMap::new);
if let Some(client) = CLIENTS.get(key) {
return Ok(client.clone());
}
let built = build()?;
// Two first calls may race to build; both get the one that landed in the map.
Ok(CLIENTS.entry(key.to_string()).or_insert(built).clone())
}

/// Validate a URL for SSRF (Server-Side Request Forgery) protection.
///
/// `allowed_schemes` must be lowercase (the URL's scheme is lowercased before
Expand All @@ -111,6 +134,21 @@ pub fn validate_url_for_ssrf(url: &str, allowed_schemes: &[&str]) -> Result<(),
validate_url_for_ssrf_inner(url, allowed_schemes, loopback_endpoints_allowed())
}

/// [`validate_url_for_ssrf`] without the resolve-then-validate step: the scheme allowlist, the
/// blocked hostnames and every IP-literal spelling, and nothing that does I/O.
///
/// For a synchronous constructor on the request path. `validate_url_for_ssrf` resolves a DNS
/// name with a blocking `getaddrinfo`, and a constructor that runs on a tokio worker holds that
/// worker for the whole lookup (~70 ms for an external name under Kubernetes' `ndots:5`), so
/// every other session on the replica waits with it. Pair this with the full check off the
/// workers (`tokio::task::spawn_blocking`) before the first dial.
pub fn validate_url_for_ssrf_without_dns(
url: &str,
allowed_schemes: &[&str],
) -> Result<(), String> {
ssrf_dns_host(url, allowed_schemes, loopback_endpoints_allowed()).map(|_| ())
}

/// Build a reqwest redirect policy that validates every redirect target before
/// following it. Use this for requests whose original URL passed
/// [`validate_url_for_ssrf`]; reqwest follows redirects by default, and an
Expand Down Expand Up @@ -158,6 +196,23 @@ fn validate_url_for_ssrf_inner(
allowed_schemes: &[&str],
loopback_allowed: bool,
) -> Result<(), String> {
match ssrf_dns_host(url, allowed_schemes, loopback_allowed)? {
// Resolve-then-validate: when the host is a DNS name (not an IP literal),
// resolve it and reject if ANY resolved address is private/internal. This
// closes DNS-rebinding / TOCTOU holes where a public-looking hostname
// resolves to a private/metadata address.
Some(host) => validate_resolved_host_for_ssrf(&host),
None => Ok(()),
}
}

/// Every check [`validate_url_for_ssrf_inner`] makes that needs no I/O. `Ok(Some(host))` when
/// the host is a DNS name still to be resolved, `Ok(None)` when nothing is left to check.
fn ssrf_dns_host(
url: &str,
allowed_schemes: &[&str],
loopback_allowed: bool,
) -> Result<Option<String>, String> {
let parsed = url::Url::parse(url).map_err(|e| format!("invalid URL '{}': {}", url, e))?;

// Scheme allowlist — applies even when the loopback escape hatch is on.
Expand All @@ -172,7 +227,7 @@ fn validate_url_for_ssrf_inner(

// Test/local-mock escape hatch (opt-in, OFF by default).
if loopback_allowed {
return Ok(());
return Ok(None);
}

let host = parsed
Expand All @@ -194,7 +249,7 @@ fn validate_url_for_ssrf_inner(
ip
));
}
return Ok(());
return Ok(None);
}

// Bracketed IPv6 literal.
Expand All @@ -208,7 +263,7 @@ fn validate_url_for_ssrf_inner(
));
}
// An IP literal — never DNS-resolved.
return Ok(());
return Ok(None);
}

// DECIMAL/integer IPv4 literal (e.g. `http://3232235777` == 192.168.1.1).
Expand All @@ -225,14 +280,10 @@ fn validate_url_for_ssrf_inner(
host, ip
));
}
return Ok(());
return Ok(None);
}

// Resolve-then-validate: when the host is a DNS name (not an IP literal),
// resolve it and reject if ANY resolved address is private/internal. This
// closes DNS-rebinding / TOCTOU holes where a public-looking hostname
// resolves to a private/metadata address.
validate_resolved_host_for_ssrf(host)
Ok(Some(host.to_string()))
}

/// Resolve a DNS hostname and reject if any resolved IP is private/internal.
Expand Down Expand Up @@ -383,6 +434,30 @@ pub(crate) fn ssrf_env_lock() -> std::sync::MutexGuard<'static, ()> {

#[cfg(test)]
mod tests {

#[test]
fn a_shared_client_is_built_once_per_key() {
use std::sync::atomic::{AtomicUsize, Ordering};
let builds = AtomicUsize::new(0);
let build = || {
builds.fetch_add(1, Ordering::SeqCst);
reqwest::Client::builder().http1_only().build()
};
let key = "net-tests-shared-client-once";
shared_http_client(key, build).unwrap();
shared_http_client(key, build).unwrap();
assert_eq!(
builds.load(Ordering::SeqCst),
1,
"the second call reuses the first client"
);
shared_http_client("net-tests-shared-client-other", build).unwrap();
assert_eq!(
builds.load(Ordering::SeqCst),
2,
"another key builds its own"
);
}
use super::*;

const HTTP_SCHEMES: &[&str] = &["http", "https"];
Expand Down Expand Up @@ -544,6 +619,36 @@ mod tests {
);
}

/// The DNS-free variant makes every check but the resolution: a name that resolves to a
/// private address passes it (the full check refuses it), everything else is refused alike.
#[test]
fn without_dns_skips_only_the_resolution() {
let _guard = env_guard();
for bad in [
"ftp://example.com/x",
"https://127.0.0.1/x",
"https://10.0.0.5/x",
"https://[::1]/x",
"https://3232235777/x",
"https://localhost/x",
"https://169.254.169.254/x",
] {
assert!(
validate_url_for_ssrf_without_dns(bad, HTTP_SCHEMES).is_err(),
"{bad}"
);
assert!(validate_url_for_ssrf(bad, HTTP_SCHEMES).is_err(), "{bad}");
}
assert!(validate_url_for_ssrf_without_dns("https://8.8.8.8/x", HTTP_SCHEMES).is_ok());
// `localhost.` (trailing dot) is not on the blocked list but resolves to loopback.
if validate_resolved_host_for_ssrf("localhost.").is_err() {
assert!(
validate_url_for_ssrf_without_dns("https://localhost./x", HTTP_SCHEMES).is_ok()
);
assert!(validate_url_for_ssrf("https://localhost./x", HTTP_SCHEMES).is_err());
}
}

/// Loopback gate OFF (default): private targets rejected via the public,
/// env-reading entry point (env var removed under the shared lock).
#[test]
Expand Down
69 changes: 68 additions & 1 deletion gateway/src/core/state.rs
Original file line number Diff line number Diff line change
Expand Up @@ -22,7 +22,9 @@ use crate::core::turn_detect::TurnDetector;
#[cfg(feature = "turn-detect")]
use crate::core::turn_detect::{TurnDetector, TurnDetectorConfig};
use crate::state::SipHooksState;
use crate::utils::req_manager::ReqManager;
use crate::utils::req_manager::{
DeploymentReqManagers, MAX_CONCURRENT_REQUESTS, ReqManager, ReqManagerConfig,
};

/// Core-specific shared state for the application.
///
Expand All @@ -32,6 +34,10 @@ use crate::utils::req_manager::ReqManager;
pub struct CoreState {
/// HTTP request managers for TTS providers - key is provider name (e.g., "deepgram")
pub tts_req_managers: Arc<RwLock<HashMap<String, Arc<ReqManager>>>>,
/// One pooled request manager per deployment for one-shot synthesis (`/v1/audio/speech`)
/// with a vendor that has no shared per-vendor manager above -- see
/// [`CoreState::deployment_tts_req_manager`].
pub deployment_tts_req_managers: Arc<DeploymentReqManagers>,
/// Unified cache store (in-memory by default)
pub cache: Arc<CacheStore>,
/// Turn detector for determining end of user speech turns
Expand Down Expand Up @@ -160,6 +166,7 @@ impl CoreState {

Ok(Arc::new(Self {
tts_req_managers: Arc::new(RwLock::new(tts_req_managers)),
deployment_tts_req_managers: Arc::new(DeploymentReqManagers::new()),
cache,
turn_detector,
sip_hooks_state,
Expand All @@ -185,6 +192,53 @@ impl CoreState {
self.tts_req_managers.read().await.get(provider).cloned()
}

/// The pooled manager for one-shot synthesis against this deployment's endpoint, built on
/// first use. `None` only if the manager cannot be built, in which case the provider falls
/// back to building its own as before.
///
/// Keyed by vendor, endpoint base and the transport timeouts -- what makes two deployments
/// need different connections. The credential is not part of it: each request carries its
/// own. The permit count is [`tts_max_concurrent_per_deployment`], and like the per-vendor
/// knob it bounds transport concurrency on this replica only; vendor-ACCOUNT concurrency is
/// the deployment's `max_concurrent`.
pub async fn deployment_tts_req_manager(
&self,
config: &crate::core::tts::TTSConfig,
) -> Option<Arc<ReqManager>> {
let key = format!(
"{}|{}|{:?}|{:?}",
config.provider.trim().to_ascii_lowercase(),
config
.api_base
.as_deref()
.map(str::trim)
.unwrap_or_default(),
config.connection_timeout,
config.request_timeout,
);
let mut req_config = ReqManagerConfig {
max_concurrent_requests: tts_max_concurrent_per_deployment(),
..Default::default()
};
if let Some(secs) = config.connection_timeout {
req_config.connect_timeout = std::time::Duration::from_secs(secs);
}
if let Some(secs) = config.request_timeout {
req_config.request_timeout = std::time::Duration::from_secs(secs);
}
match self
.deployment_tts_req_managers
.get_or_create(&key, req_config)
.await
{
Ok(manager) => Some(manager),
Err(e) => {
tracing::warn!(error = %e, "could not build the per-deployment TTS request manager");
None
}
}
}

#[cfg(feature = "turn-detect")]
/// Initialize and warmup the Turn Detector model
async fn initialize_turn_detector(
Expand Down Expand Up @@ -348,6 +402,19 @@ fn parse_env_positive_usize(name: &str) -> Result<Option<usize>, String> {

/// `WAAV_TTS_MAX_CONCURRENT_PER_VENDOR` (default 64, 1–1000): concurrent vendor requests per TTS
/// vendor per replica. It was a hard-coded 4 with an unbounded queue behind it.
/// Permits of each per-deployment one-shot TTS manager
/// ([`CoreState::deployment_tts_req_manager`]): `WAAV_TTS_MAX_CONCURRENT_PER_DEPLOYMENT`,
/// 1..=[`MAX_CONCURRENT_REQUESTS`], default 4096. It is shared by every speech request to the
/// deployment on this replica, so it has to hold the replica's whole load for that deployment:
/// at 5 s of vendor time per request, 4096 is ~800 requests a second before the bounded wait.
pub fn tts_max_concurrent_per_deployment() -> usize {
std::env::var("WAAV_TTS_MAX_CONCURRENT_PER_DEPLOYMENT")
.ok()
.and_then(|v| v.trim().parse::<usize>().ok())
.filter(|n| (1..=MAX_CONCURRENT_REQUESTS).contains(n))
.unwrap_or(4096)
}

pub fn tts_max_concurrent_per_vendor() -> usize {
std::env::var("WAAV_TTS_MAX_CONCURRENT_PER_VENDOR")
.ok()
Expand Down
14 changes: 9 additions & 5 deletions gateway/src/core/stt/prerecorded.rs
Original file line number Diff line number Diff line change
Expand Up @@ -172,12 +172,16 @@ type AsyncErrorCallback = Box<
+ Sync,
>;

/// Pooled and shared by every prerecorded upload (`core::net::shared_http_client`): a client per
/// request threw its connections away, so each call paid DNS + TCP + TLS to the vendor.
fn http_client() -> Result<Client, reqwest::Error> {
crate::core::net::ssrf_protected_client_builder(crate::core::net::HTTP_URL_SCHEMES)
.timeout(HTTP_TIMEOUT)
.pool_max_idle_per_host(4)
.pool_idle_timeout(Duration::from_secs(90))
.build()
crate::core::net::shared_http_client("stt-prerecorded", || {
crate::core::net::ssrf_protected_client_builder(crate::core::net::HTTP_URL_SCHEMES)
.timeout(HTTP_TIMEOUT)
.pool_idle_timeout(Duration::from_secs(90))
.http1_only()
.build()
})
}

/// Buffers PCM and transcribes it against a vendor's prerecorded API on close.
Expand Down
Loading
Loading