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
}
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: &str) {
counter!(SUBREQUEST_STREAMS_TOTAL, "termination" => termination.to_owned()).increment(1);
histogram!(SUBREQUEST_STREAM_DURATION_SECONDS).record(0.0);
debug!(termination, "sub-request: stream terminated at header phase");
}
pub(super) const HOP_BY_HOP_HEADERS: &[&str] = &[
"connection",
"keep-alive",
"proxy-authenticate",
"proxy-authorization",
"te",
"trailer",
"transfer-encoding",
"upgrade",
];
pub(super) fn strip_hop_by_hop_headers(headers: &mut HeaderMap) {
let connection_values: Vec<_> = headers.get_all(http::header::CONNECTION).iter().cloned().collect();
for name in HOP_BY_HOP_HEADERS {
headers.remove(*name);
}
for value in connection_values {
let Ok(value) = value.to_str() else { continue };
for token in value.split(',').map(str::trim).filter(|token| !token.is_empty()) {
headers.remove(token);
}
}
}
pub(super) fn strip_request_framing_headers(headers: &mut HeaderMap) {
headers.remove(http::header::CONTENT_LENGTH);
headers.remove(http::header::TRANSFER_ENCODING);
}
pub(super) fn strip_reserved_headers(headers: &mut HeaderMap) {
let reserved: Vec<http::header::HeaderName> = headers
.keys()
.filter(|name| crate::reserved_headers::is_reserved(name.as_str()))
.cloned()
.collect();
for name in reserved {
headers.remove(&name);
}
}
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"))
}
}