use std::{sync::Arc, time::Duration};
use http::HeaderMap;
use metrics::{counter, histogram};
use pingora_core::{
connectors::{ConnectorOptions, http::Connector},
protocols::http::client::HttpSession,
upstreams::peer::{HttpPeer, Peer as _},
};
use tokio::sync::{OwnedSemaphorePermit, Semaphore};
use tracing::debug;
use super::types::{SubRequestError, SubResponse};
use crate::circuit::{CircuitBreakerConfig, CircuitBreakerRegistry, CircuitToken, PeerKey};
pub(super) const SUBREQUEST_STREAMS_TOTAL: &str = "praxis_subrequest_streams_total";
pub(super) const SUBREQUEST_STREAM_DURATION_SECONDS: &str = "praxis_subrequest_stream_duration_seconds";
pub(super) const SUBREQUEST_STREAM_BYTES_TOTAL: &str = "praxis_subrequest_stream_bytes_total";
pub(super) const SUBREQUEST_HEADER_DURATION_SECONDS: &str = "praxis_subrequest_header_duration_seconds";
#[derive(Debug)]
pub struct SubRequestConnectorOptions {
pub keepalive_pool_size: usize,
pub max_connections: Option<usize>,
pub circuit_breaker: Option<CircuitBreakerConfig>,
}
#[derive(Clone)]
pub struct SubRequestConnector {
pub(super) inner: Arc<Connector<()>>,
pub(super) admission: Option<Arc<Semaphore>>,
pub(super) configured_max_connections: Option<usize>,
pub(super) circuit_breakers: Option<Arc<CircuitBreakerRegistry>>,
}
impl SubRequestConnector {
pub fn new(keepalive_pool_size: usize, max_connections: Option<usize>) -> Self {
let options = ConnectorOptions::new(keepalive_pool_size);
Self {
inner: Arc::new(Connector::new(Some(options))),
admission: max_connections.map(|n| Arc::new(Semaphore::new(n))),
configured_max_connections: max_connections,
circuit_breakers: None,
}
}
pub fn with_options(opts: SubRequestConnectorOptions) -> Self {
let options = ConnectorOptions::new(opts.keepalive_pool_size);
Self {
inner: Arc::new(Connector::new(Some(options))),
admission: opts.max_connections.map(|n| Arc::new(Semaphore::new(n))),
configured_max_connections: opts.max_connections,
circuit_breakers: opts
.circuit_breaker
.map(|cfg| Arc::new(CircuitBreakerRegistry::new(cfg))),
}
}
pub fn connector(&self) -> &Connector<()> {
&self.inner
}
#[must_use]
pub fn has_circuit_breaker(&self) -> bool {
self.circuit_breakers.is_some()
}
#[must_use]
pub fn configured_max_connections(&self) -> Option<usize> {
self.configured_max_connections
}
pub async fn acquire_permit(&self) -> Option<OwnedSemaphorePermit> {
let semaphore = self.admission.as_ref()?;
Arc::clone(semaphore).acquire_owned().await.ok()
}
pub async fn try_acquire_permit(&self, timeout: Duration) -> Result<Option<OwnedSemaphorePermit>, SubRequestError> {
let Some(semaphore) = self.admission.as_ref() else {
return Ok(None);
};
let configured = self.configured_max_connections.unwrap_or(0);
match tokio::time::timeout(timeout, Arc::clone(semaphore).acquire_owned()).await {
Ok(Ok(permit)) => Ok(Some(permit)),
Ok(Err(_closed)) => Err(SubRequestError::AdmissionTimeout {
max_connections: configured,
}),
Err(_elapsed) => Err(SubRequestError::AdmissionTimeout {
max_connections: configured,
}),
}
}
}
impl std::fmt::Debug for SubRequestConnector {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SubRequestConnector")
.field("pool", &"Connector<()>")
.field("max_connections", &self.configured_max_connections)
.field("circuit_breakers", &self.circuit_breakers.is_some())
.finish()
}
}
pub(super) struct RawExchange<'a> {
pub(super) session: HttpSession<()>,
pub(super) peer: HttpPeer,
pub(super) connector: &'a SubRequestConnector,
pub(super) status: u16,
pub(super) headers: HeaderMap,
pub(super) circuit_guard: Option<CircuitGuard<'a>>,
pub(super) permit: Option<OwnedSemaphorePermit>,
pub(super) deadline: tokio::time::Instant,
}
pub(super) struct CircuitGuard<'a> {
registry: &'a CircuitBreakerRegistry,
peer: PeerKey,
token: Option<CircuitToken>,
}
impl<'a> CircuitGuard<'a> {
pub(super) fn new(registry: &'a CircuitBreakerRegistry, peer: PeerKey, token: CircuitToken) -> Self {
Self {
registry,
peer,
token: Some(token),
}
}
pub(super) fn finalize_success(mut self) {
if let Some(token) = self.token.take() {
self.registry.record_success(&self.peer, token);
}
}
pub(super) fn finalize(mut self, result: &Result<SubResponse, SubRequestError>) {
let Some(token) = self.token.take() else {
return;
};
match result {
Err(SubRequestError::Connect(_) | SubRequestError::Io(_) | SubRequestError::DeadlineExceeded) => {
self.registry.record_failure(&self.peer, token);
},
Ok(_) | Err(_) => {
self.registry.record_success(&self.peer, token);
},
}
}
}
impl Drop for CircuitGuard<'_> {
fn drop(&mut self) {
if let Some(token) = self.token.take() {
self.registry.record_failure(&self.peer, token);
}
}
}
pub(super) fn check_clean_completion(session: &mut HttpSession<()>) -> Result<bool, SubRequestError> {
use pingora_core::protocols::http::custom::client::Session as _;
match session {
HttpSession::H1(h1) => Ok(h1.is_body_done()),
HttpSession::H2(h2) => h2
.check_response_end_or_error()
.map_err(|e| SubRequestError::Io(e.to_string())),
HttpSession::Custom(c) => Ok(c.response_finished()),
}
}
pub(super) fn record_header_termination(termination: &'static str) {
counter!(SUBREQUEST_STREAMS_TOTAL, "termination" => termination).increment(1);
histogram!(SUBREQUEST_STREAM_DURATION_SECONDS).record(0.0);
debug!(termination, "sub-request: stream terminated at header phase");
}
pub(super) use crate::reserved_headers::HOP_BY_HOP_HEADERS;
pub(super) fn connection_nominated_tokens(headers: &HeaderMap) -> Vec<&str> {
headers
.get_all(http::header::CONNECTION)
.iter()
.filter_map(|value| value.to_str().ok())
.flat_map(|value| value.split(','))
.map(str::trim)
.filter(|token| !token.is_empty())
.collect()
}
pub(super) fn is_boundary_stripped(name: &http::header::HeaderName, nominated: &[&str]) -> bool {
let name = name.as_str();
HOP_BY_HOP_HEADERS.contains(&name)
|| crate::reserved_headers::is_reserved(name)
|| (nominated.iter().any(|token| token.eq_ignore_ascii_case(name))
&& !crate::reserved_headers::is_connection_token_protected(name))
}
pub(super) fn is_request_stripped(name: &http::header::HeaderName, nominated: &[&str]) -> bool {
let lower = name.as_str();
lower == "content-length" || lower == "transfer-encoding" || is_boundary_stripped(name, nominated)
}
pub(super) fn is_transport_header(name: &http::header::HeaderName) -> bool {
HOP_BY_HOP_HEADERS.iter().any(|h| *h == name.as_str()) || name == http::header::CONTENT_LENGTH
}
pub(super) fn empty_body_needs_framing(method: &http::Method) -> bool {
matches!(*method, http::Method::POST | http::Method::PUT | http::Method::PATCH)
}
pub(super) fn ensure_host_header(
request: &mut pingora_http::RequestHeader,
peer: &HttpPeer,
) -> Result<(), SubRequestError> {
if !request.headers.contains_key(http::header::HOST) {
request
.insert_header(http::header::HOST, peer.address().to_string())
.map_err(|error| SubRequestError::InvalidRequest(error.to_string()))?;
}
Ok(())
}
pub(super) fn clamp_peer_timeouts(peer: &mut HttpPeer, deadline: Duration) {
peer.options.connection_timeout = Some(min_timeout(peer.options.connection_timeout, deadline));
peer.options.total_connection_timeout = Some(min_timeout(peer.options.total_connection_timeout, deadline));
}
pub(super) fn min_timeout(configured: Option<Duration>, deadline: Duration) -> Duration {
configured.map_or(deadline, |configured| configured.min(deadline))
}
pub(super) fn classify_timeout(
remaining_budget: Duration,
configured_timeout: Option<Duration>,
phase: &str,
) -> SubRequestError {
if configured_timeout.is_none_or(|t| t >= remaining_budget) {
SubRequestError::DeadlineExceeded
} else {
SubRequestError::Io(format!("upstream {phase} timeout"))
}
}