use async_trait::async_trait;
use bytes::Bytes;
use futures::future::BoxFuture;
use futures::future::FutureExt;
use http::{header, version::Version, Method};
use log::{debug, error, trace, warn};
use once_cell::sync::Lazy;
use pingora_http::{RequestHeader, ResponseHeader};
use std::fmt::Debug;
use std::future::{poll_fn, Future};
use std::str;
use std::sync::{
atomic::{AtomicBool, AtomicU64, AtomicU8, AtomicUsize, Ordering},
Arc,
};
use std::task::Poll;
use std::time::Duration;
use tokio::sync::{mpsc, Notify};
use tokio::time;
use pingora_cache::NoCacheReason;
use pingora_core::apps::{
HttpPersistentSettings, HttpServerApp, HttpServerOptions, ReusedHttpStream,
};
use pingora_core::connectors::http::custom;
use pingora_core::connectors::{http::Connector, ConnectorOptions};
use pingora_core::modules::http::compression::ResponseCompressionBuilder;
use pingora_core::modules::http::{HttpModuleCtx, HttpModules};
use pingora_core::protocols::http::client::HttpSession as ClientSession;
use pingora_core::protocols::http::custom::CustomMessageWrite;
use pingora_core::protocols::http::subrequest::server::SubrequestHandle;
use pingora_core::protocols::http::v1::client::HttpSession as HttpSessionV1;
use pingora_core::protocols::http::v2::server::H2Options;
use pingora_core::protocols::http::HttpTask;
use pingora_core::protocols::http::ServerSession as HttpSession;
use pingora_core::protocols::http::SERVER_NAME;
use pingora_core::protocols::Stream;
use pingora_core::protocols::{Digest, UniqueID};
use pingora_core::server::configuration::ServerConf;
use pingora_core::server::{RuntimeOpts, ShutdownWatch};
use pingora_core::upstreams::peer::{HttpPeer, Peer};
use pingora_error::{Error, ErrorSource, ErrorType::*, OrErr, Result};
const TASK_BUFFER_SIZE: usize = 4;
const MAX_SHUTDOWN_NOTIFY_SHARDS: usize = 256;
type DownstreamCustomMessageReader =
Box<dyn futures::Stream<Item = Result<Bytes>> + Unpin + Send + Sync + 'static>;
mod proxy_cache;
mod proxy_common;
mod proxy_custom;
mod proxy_h1;
mod proxy_h2;
mod proxy_purge;
mod proxy_trait;
pub mod subrequest;
use subrequest::{BodyMode, Ctx as SubrequestCtx};
pub use proxy_cache::range_filter::{range_header_filter, MultiRangeInfo, RangeType};
pub use proxy_purge::PurgeStatus;
pub use proxy_trait::{FailToProxy, ProxyHttp, ProxyWarnLogContext};
pub mod prelude {
pub use crate::{http_proxy, http_proxy_service, ProxyHttp, ProxyWarnLogContext, Session};
}
pub type ProcessCustomSession<SV, C> = Arc<
dyn Fn(Arc<HttpProxy<SV, C>>, Stream, &ShutdownWatch) -> BoxFuture<'static, Option<Stream>>
+ Send
+ Sync
+ Unpin
+ 'static,
>;
struct ShardedNotify {
shards: Box<[NotifyShard]>,
}
#[repr(align(128))]
struct NotifyShard(Notify);
impl ShardedNotify {
fn new(worker_threads: usize) -> Self {
let shards = worker_threads
.max(1)
.checked_next_power_of_two()
.unwrap_or(MAX_SHUTDOWN_NOTIFY_SHARDS)
.min(MAX_SHUTDOWN_NOTIFY_SHARDS);
ShardedNotify {
shards: (0..shards).map(|_| NotifyShard(Notify::new())).collect(),
}
}
fn local(&self) -> &Notify {
static NEXT_THREAD_ID: AtomicUsize = AtomicUsize::new(0);
thread_local! {
static THREAD_ID: usize = NEXT_THREAD_ID.fetch_add(1, Ordering::Relaxed);
}
let id = THREAD_ID.with(|id| *id);
&self.shards[id & (self.shards.len() - 1)].0
}
fn notify_waiters(&self) {
for shard in self.shards.iter() {
shard.0.notify_waiters();
}
}
}
pub struct HttpProxy<SV, C = ()>
where
C: custom::Connector, {
inner: SV, client_upstream: Connector<C>,
shutdown: ShardedNotify,
shutdown_flag: Arc<AtomicBool>,
pub server_options: Option<HttpServerOptions>,
pub h2_options: Option<H2Options>,
pub downstream_modules: HttpModules,
#[cfg(feature = "upstream_modules")]
pub upstream_modules: HttpModules,
max_retries: usize,
process_custom_session: Option<ProcessCustomSession<SV, C>>,
}
impl<SV> HttpProxy<SV, ()> {
pub fn new(inner: SV, conf: Arc<ServerConf>) -> Self {
HttpProxy {
inner,
client_upstream: Connector::new(Some(ConnectorOptions::from_server_conf(&conf))),
shutdown: ShardedNotify::new(conf.threads),
shutdown_flag: Arc::new(AtomicBool::new(false)),
server_options: None,
h2_options: None,
downstream_modules: HttpModules::new(),
#[cfg(feature = "upstream_modules")]
upstream_modules: HttpModules::new(),
max_retries: conf.max_retries,
process_custom_session: None,
}
}
}
impl<SV, C> HttpProxy<SV, C>
where
C: custom::Connector,
{
fn new_custom(
inner: SV,
conf: Arc<ServerConf>,
connector: C,
on_custom: Option<ProcessCustomSession<SV, C>>,
server_options: Option<HttpServerOptions>,
client_options: Option<ConnectorOptions>,
) -> Self
where
SV: ProxyHttp + Send + Sync + 'static,
SV::CTX: Send + Sync,
{
let client_options =
client_options.unwrap_or_else(|| ConnectorOptions::from_server_conf(&conf));
let client_upstream = Connector::new_custom(Some(client_options), connector);
HttpProxy {
inner,
client_upstream,
shutdown: ShardedNotify::new(conf.threads),
shutdown_flag: Arc::new(AtomicBool::new(false)),
server_options,
downstream_modules: HttpModules::new(),
#[cfg(feature = "upstream_modules")]
upstream_modules: HttpModules::new(),
max_retries: conf.max_retries,
process_custom_session: on_custom,
h2_options: None,
}
}
pub fn unexpected_data_connection_count(&self) -> u64 {
self.client_upstream.unexpected_data_connection_count()
}
pub fn unexpected_data_connection_counter(&self) -> Arc<AtomicU64> {
self.client_upstream.unexpected_data_connection_counter()
}
pub fn handle_init_modules(&mut self)
where
SV: ProxyHttp,
{
self.inner
.init_downstream_modules(&mut self.downstream_modules);
#[cfg(feature = "upstream_modules")]
self.inner.init_upstream_modules(&mut self.upstream_modules);
}
async fn await_shutdown(&self) {
let notified = self.shutdown.local().notified();
tokio::pin!(notified);
poll_fn(|context| {
if notified.as_mut().poll(context).is_ready()
|| self.shutdown_flag.load(Ordering::Acquire)
{
Poll::Ready(())
} else {
Poll::Pending
}
})
.await;
}
async fn handle_new_request(
&self,
mut downstream_session: Box<HttpSession>,
) -> Option<Box<HttpSession>>
where
SV: ProxyHttp + Send + Sync,
SV::CTX: Send + Sync,
{
let res = tokio::select! {
biased; res = downstream_session.read_request() => { res }
_ = self.await_shutdown() => {
return None;
}
};
match res {
Ok(true) => {
debug!("Successfully get a new request");
}
Ok(false) => {
return None; }
Err(mut e) => {
e.as_down();
if matches!(e.etype, InvalidHTTPHeader) {
debug!(
"Fail to proxy: {e}, downstream session type: {}",
downstream_session.session_type()
);
downstream_session
.respond_error(400)
.await
.unwrap_or_else(|e| {
error!("failed to send error response to downstream: {e}");
});
} else {
error!(
"Fail to proxy: {e}, downstream session type: {}",
downstream_session.session_type()
);
}
downstream_session.shutdown().await;
return None;
}
}
trace!(
"Request header: {:?}",
downstream_session.req_header().as_ref()
);
if !self
.server_options
.as_ref()
.is_some_and(|opts| opts.allow_connect_method_proxying)
&& downstream_session.req_header().method == Method::CONNECT
{
downstream_session
.respond_error(405)
.await
.unwrap_or_else(|e| {
error!("failed to send error response to downstream: {e}");
});
downstream_session.shutdown().await;
return None;
}
Some(downstream_session)
}
async fn proxy_to_upstream(
&self,
session: &mut Session,
ctx: &mut SV::CTX,
) -> (bool, Option<Box<Error>>)
where
SV: ProxyHttp + Send + Sync,
SV::CTX: Send + Sync,
{
let peer = match self.inner.upstream_peer(session, ctx).await {
Ok(p) => p,
Err(e) => return (false, Some(e)),
};
let client_session = self.client_upstream.get_http_session(&*peer).await;
match client_session {
Ok((client_session, client_reused)) => {
let (server_reused, error) = match client_session {
ClientSession::H1(mut h1) => {
let (server_reused, client_reuse, error) = self
.proxy_to_h1_upstream(session, &mut h1, client_reused, &peer, ctx)
.await;
if client_reuse {
let session = ClientSession::H1(h1);
self.client_upstream
.release_http_session(session, &*peer, peer.idle_timeout())
.await;
}
(server_reused, error)
}
ClientSession::H2(mut h2) => {
let (server_reused, mut error) = self
.proxy_to_h2_upstream(session, &mut h2, client_reused, &peer, ctx)
.await;
let session = ClientSession::H2(h2);
self.client_upstream
.release_http_session(session, &*peer, peer.idle_timeout())
.await;
if let Some(e) = error.as_mut() {
if matches!(e.etype, H2Downgrade | InvalidH2) {
if peer
.get_alpn()
.is_none_or(|alpn| alpn.get_min_http_version() == 1)
{
self.client_upstream.prefer_h1(&*peer);
} else {
e.retry = false.into();
}
}
}
(server_reused, error)
}
ClientSession::Custom(mut c) => {
let (server_reused, error) = self
.proxy_to_custom_upstream(session, &mut c, client_reused, &peer, ctx)
.await;
let session = ClientSession::Custom(c);
self.client_upstream
.release_http_session(session, &*peer, peer.idle_timeout())
.await;
(server_reused, error)
}
};
(
server_reused,
error.map(|e| {
self.inner
.error_while_proxy(&peer, session, e, ctx, client_reused)
}),
)
}
Err(mut e) => {
e.as_up();
let new_err = self.inner.fail_to_connect(session, &peer, ctx, e);
(false, Some(new_err.into_up()))
}
}
}
async fn upstream_filter(
&self,
session: &mut Session,
task: &mut HttpTask,
ctx: &mut SV::CTX,
) -> Result<Option<Duration>>
where
SV: ProxyHttp + Send + Sync,
SV::CTX: Send + Sync,
{
let duration = match task {
HttpTask::Header(header, _eos) => {
self.inner
.upstream_response_filter(session, header, ctx)
.await?;
None
}
HttpTask::Body(data, eos) | HttpTask::UpgradedBody(data, eos) => self
.inner
.upstream_response_body_filter(session, data, *eos, ctx)?,
HttpTask::Trailer(Some(trailers)) => {
self.inner
.upstream_response_trailer_filter(session, trailers, ctx)?;
None
}
_ => {
None
}
};
Ok(duration)
}
async fn finish(
&self,
mut session: Session,
ctx: &mut SV::CTX,
reuse: bool,
error: Option<Box<Error>>,
) -> Option<ReusedHttpStream>
where
SV: ProxyHttp + Send + Sync,
SV::CTX: Send + Sync,
{
self.inner
.logging(&mut session, error.as_deref(), ctx)
.await;
if let Some(e) = error {
session.downstream_session.on_proxy_failure(e);
}
if reuse {
let mut persistent_settings = HttpPersistentSettings::for_session(&session);
if let Some(uc) = self.inner.persist_connection_context(&session, ctx) {
persistent_settings.set_user_context(uc);
}
session
.downstream_session
.finish()
.await
.ok()
.flatten()
.map(|s| ReusedHttpStream::from_reusable_stream(s, persistent_settings))
} else {
None
}
}
fn cleanup_sub_req(&self, session: &mut Session) {
if let Some(ctx) = session.subrequest_ctx.as_mut() {
ctx.release_write_lock();
}
}
}
use pingora_cache::HttpCache;
use pingora_core::protocols::http::compression::ResponseCompressionCtx;
pub struct Session {
pub downstream_session: Box<HttpSession>,
pub cache: HttpCache,
pub upstream_compression: ResponseCompressionCtx,
pub ignore_downstream_range: bool,
pub upstream_headers_mutated_for_cache: bool,
h1_upgrade_request_status: H1UpgradeRequestStatus,
pub subrequest_ctx: Option<Box<SubrequestCtx>>,
pub subrequest_spawner: Option<SubrequestSpawner>,
pub downstream_modules_ctx: HttpModuleCtx,
#[cfg(feature = "upstream_modules")]
pub upstream_modules_ctx: HttpModuleCtx,
upstream_body_bytes_received: usize,
upstream_body_bytes_sent: Option<usize>,
downstream_task_seen_upgraded: bool,
upstream_write_pending_time: Duration,
shutdown_flag: Arc<AtomicBool>,
}
impl Session {
fn new(
downstream_session: impl Into<Box<HttpSession>>,
downstream_modules: &HttpModules,
#[cfg(feature = "upstream_modules")] upstream_modules: &HttpModules,
shutdown_flag: Arc<AtomicBool>,
) -> Self {
Session {
downstream_session: downstream_session.into(),
cache: HttpCache::new(),
upstream_compression: ResponseCompressionCtx::new(0, false, false),
ignore_downstream_range: false,
upstream_headers_mutated_for_cache: false,
h1_upgrade_request_status: H1UpgradeRequestStatus::default(),
subrequest_ctx: None,
subrequest_spawner: None, downstream_modules_ctx: downstream_modules.build_ctx(),
#[cfg(feature = "upstream_modules")]
upstream_modules_ctx: upstream_modules.build_ctx(),
upstream_body_bytes_received: 0,
upstream_body_bytes_sent: None,
downstream_task_seen_upgraded: false,
upstream_write_pending_time: Duration::ZERO,
shutdown_flag,
}
}
pub fn new_h1(stream: Stream) -> Self {
let modules = HttpModules::new();
Self::new(
Box::new(HttpSession::new_http1(stream)),
&modules,
#[cfg(feature = "upstream_modules")]
&HttpModules::new(),
Arc::new(AtomicBool::new(false)),
)
}
pub fn new_h1_with_modules(stream: Stream, downstream_modules: &HttpModules) -> Self {
Self::new(
Box::new(HttpSession::new_http1(stream)),
downstream_modules,
#[cfg(feature = "upstream_modules")]
&HttpModules::new(),
Arc::new(AtomicBool::new(false)),
)
}
#[cfg(feature = "upstream_modules")]
pub async fn upstream_modules_filter_task(&mut self, t: &mut HttpTask) -> Result<()> {
match t {
HttpTask::Header(header, eos) => {
self.upstream_modules_ctx
.response_header_filter(header, *eos)
.await?;
}
HttpTask::Body(body, eos) | HttpTask::UpgradedBody(body, eos) => {
self.upstream_modules_ctx.response_body_filter(body, *eos)?;
}
HttpTask::Trailer(trailers) => {
if let Some(buf) = self
.upstream_modules_ctx
.response_trailer_filter(trailers)?
{
*t = HttpTask::Body(Some(buf), true);
}
}
HttpTask::Done => {
if let Some(buf) = self.upstream_modules_ctx.response_done_filter()? {
*t = HttpTask::Body(Some(buf), true);
}
}
HttpTask::Failed(_) => {}
}
Ok(())
}
pub fn as_downstream_mut(&mut self) -> &mut HttpSession {
&mut self.downstream_session
}
pub fn as_downstream(&self) -> &HttpSession {
&self.downstream_session
}
pub async fn respond_error(&mut self, error: u16) -> Result<()> {
self.as_downstream_mut().respond_error(error).await
}
pub async fn respond_error_with_body(&mut self, error: u16, body: Bytes) -> Result<()> {
self.as_downstream_mut()
.respond_error_with_body(error, body)
.await
}
pub async fn write_response_header(
&mut self,
mut resp: Box<ResponseHeader>,
end_of_stream: bool,
) -> Result<()> {
self.downstream_modules_ctx
.response_header_filter(&mut resp, end_of_stream)
.await?;
self.downstream_session.write_response_header(resp).await
}
pub async fn write_response_header_ref(
&mut self,
resp: &ResponseHeader,
end_of_stream: bool,
) -> Result<(), Box<Error>> {
self.write_response_header(Box::new(resp.clone()), end_of_stream)
.await
}
pub async fn write_response_body(
&mut self,
mut body: Option<Bytes>,
end_of_stream: bool,
) -> Result<()> {
self.downstream_modules_ctx
.response_body_filter(&mut body, end_of_stream)?;
if body.is_none() && !end_of_stream {
return Ok(());
}
let data = body.unwrap_or_default();
self.downstream_session
.write_response_body(data, end_of_stream)
.await
}
async fn downstream_response_task_filter(
&mut self,
task: &mut HttpTask,
seen_upgraded: &mut bool,
) -> Result<()> {
match task {
HttpTask::Header(resp, end) => {
if *seen_upgraded {
return reject_unexpected_task_after_h1_upgrade(self, "header", *seen_upgraded);
}
self.downstream_modules_ctx
.response_header_filter(resp, *end)
.await?;
reject_mismatched_h1_upgrade_101(self, resp, "downstream_module_header_filter")
.map_err(|e| e.into_in())?;
if resp.status == http::StatusCode::SWITCHING_PROTOCOLS
&& self.downstream_session.is_upgrade(resp) == Some(true)
{
*seen_upgraded = true;
}
}
HttpTask::Body(data, end) => {
if *seen_upgraded {
return reject_unexpected_task_after_h1_upgrade(self, "body", *seen_upgraded);
}
self.downstream_modules_ctx
.response_body_filter(data, *end)?;
}
HttpTask::UpgradedBody(data, end) => {
if !*seen_upgraded {
return reject_unexpected_upgraded_body_before_h1_upgrade(self, *seen_upgraded);
}
self.downstream_modules_ctx
.response_body_filter(data, *end)?;
}
HttpTask::Trailer(trailers) => {
if *seen_upgraded {
return reject_unexpected_task_after_h1_upgrade(
self,
"trailer",
*seen_upgraded,
);
}
if let Some(buf) = self
.downstream_modules_ctx
.response_trailer_filter(trailers)?
{
*task = HttpTask::Body(Some(buf), true);
}
}
HttpTask::Done => {
if let Some(buf) = self.downstream_modules_ctx.response_done_filter()? {
*task = if *seen_upgraded {
HttpTask::UpgradedBody(Some(buf), true)
} else {
HttpTask::Body(Some(buf), true)
};
}
}
_ => { }
}
Ok(())
}
pub async fn send_downstream_proxy_task(&mut self, mut task: HttpTask) -> Result<()> {
let mut seen_upgraded = self.downstream_task_seen_upgraded || self.was_upgraded();
self.downstream_response_task_filter(&mut task, &mut seen_upgraded)
.await?;
self.downstream_task_seen_upgraded = seen_upgraded;
self.downstream_session.send_downstream_proxy_task(task);
Ok(())
}
pub fn set_proxy_tasks_enabled(&mut self, enabled: bool) {
self.downstream_session.set_proxy_tasks_enabled(enabled);
}
pub fn has_pending_downstream_tasks(&self) -> bool {
self.downstream_session.supports_proxy_task_api()
&& self.downstream_session.has_pending_downstream_proxy_tasks()
}
pub async fn write_downstream_proxy_tasks(&mut self) -> Result<bool> {
if self.downstream_session.supports_proxy_task_api() {
self.downstream_session.write_downstream_proxy_tasks().await
} else {
Ok(false)
}
}
pub async fn write_response_tasks(&mut self, mut tasks: Vec<HttpTask>) -> Result<bool> {
let mut seen_upgraded = self.downstream_task_seen_upgraded || self.was_upgraded();
for task in tasks.iter_mut() {
self.downstream_response_task_filter(task, &mut seen_upgraded)
.await?;
}
self.downstream_task_seen_upgraded = seen_upgraded;
self.downstream_session.response_duplex_vec(tasks).await
}
pub fn mark_upstream_headers_mutated_for_cache(&mut self) {
self.upstream_headers_mutated_for_cache = true;
}
pub fn upstream_headers_mutated_for_cache(&self) -> bool {
self.upstream_headers_mutated_for_cache
}
fn set_upstream_h1_upgrade_request_status(&mut self, upstream_is_upgrade_req: bool) {
self.h1_upgrade_request_status = H1UpgradeRequestStatus {
upstream: Some(upstream_is_upgrade_req),
};
}
fn h1_upgrade_request_snapshot(&self) -> H1UpgradeRequestSnapshot {
H1UpgradeRequestSnapshot {
downstream: self.downstream_session.is_upgrade_req(),
upstream: self.h1_upgrade_request_status.upstream,
}
}
pub fn upstream_body_bytes_received(&self) -> usize {
self.upstream_body_bytes_received
}
pub(crate) fn set_upstream_body_bytes_received(&mut self, n: usize) {
self.upstream_body_bytes_received = n;
}
pub fn upstream_body_bytes_sent(&self) -> Option<usize> {
self.upstream_body_bytes_sent
}
pub(crate) fn set_upstream_body_bytes_sent(&mut self, n: usize) {
self.upstream_body_bytes_sent = Some(n);
}
pub fn upstream_write_pending_time(&self) -> Duration {
self.upstream_write_pending_time
}
pub(crate) fn set_upstream_write_pending_time(&mut self, d: Duration) {
self.upstream_write_pending_time = d;
}
pub fn is_process_shutting_down(&self) -> bool {
self.shutdown_flag.load(Ordering::Acquire)
}
pub fn downstream_custom_message(&mut self) -> Result<Option<DownstreamCustomMessageReader>> {
if let Some(custom_session) = self.downstream_session.as_custom_mut() {
custom_session
.take_custom_message_reader()
.map(Some)
.ok_or(Error::explain(
ReadError,
"can't extract custom reader from downstream",
))
} else {
Ok(None)
}
}
fn take_downstream_custom_message_reader(
&mut self,
downstream_custom_message_writer: &mut Option<Box<dyn CustomMessageWrite>>,
) -> Result<Option<DownstreamCustomMessageReader>> {
if downstream_custom_message_writer.is_none() {
return Ok(None);
}
let Some(custom_session) = self.downstream_session.as_custom_mut() else {
return Ok(None);
};
let Some(reader) = custom_session.take_custom_message_reader() else {
if let Some(writer) = downstream_custom_message_writer.take() {
custom_session.restore_custom_message_writer(writer)?;
}
return Err(Error::explain(
ReadError,
"can't extract custom reader from downstream",
));
};
Ok(Some(reader))
}
}
#[derive(Clone, Copy, Debug, Default)]
struct H1UpgradeRequestStatus {
upstream: Option<bool>,
}
#[derive(Clone, Copy, Debug)]
struct H1UpgradeRequestSnapshot {
downstream: bool,
upstream: Option<bool>,
}
impl H1UpgradeRequestSnapshot {
fn mismatch(self) -> bool {
matches!(self.upstream, Some(upstream) if self.downstream != upstream)
}
}
fn reject_mismatched_h1_upgrade_101(
session: &Session,
header: &ResponseHeader,
stage: &'static str,
) -> Result<()> {
if header.status != http::StatusCode::SWITCHING_PROTOCOLS {
return Ok(());
}
let status = session.h1_upgrade_request_snapshot();
if status.mismatch() {
return Error::e_explain(
InvalidHTTPHeader,
format!(
"received 101 response with mismatched upstream/downstream upgrade status: stage={stage}, downstream_upgrade_req={}, upstream_upgrade_req={:?}, downstream_was_upgraded={}, downstream_task_seen_upgraded={}, response_version={:?}, response_upgrade_header_present={}, response_connection_header_present={}",
status.downstream,
status.upstream,
session.was_upgraded(),
session.downstream_task_seen_upgraded,
header.version,
header.headers.get(http::header::UPGRADE).is_some(),
header.headers.get(http::header::CONNECTION).is_some(),
),
);
}
Ok(())
}
fn reject_unexpected_task_after_h1_upgrade(
session: &Session,
task: &'static str,
task_filter_seen_upgraded: bool,
) -> Result<()> {
let status = session.h1_upgrade_request_snapshot();
Error::e_explain(
InvalidHTTPHeader,
format!(
"received {task} task after downstream 101 upgrade: downstream_upgrade_req={}, upstream_upgrade_req={:?}, downstream_was_upgraded={}, downstream_task_seen_upgraded={}, task_filter_seen_upgraded={}",
status.downstream,
status.upstream,
session.was_upgraded(),
session.downstream_task_seen_upgraded,
task_filter_seen_upgraded
),
)
.map_err(|e| e.into_in())
}
fn reject_unexpected_upgraded_body_before_h1_upgrade(
session: &Session,
task_filter_seen_upgraded: bool,
) -> Result<()> {
let status = session.h1_upgrade_request_snapshot();
Error::e_explain(
InvalidHTTPHeader,
format!(
"received upgraded body task before downstream 101 upgrade: downstream_upgrade_req={}, upstream_upgrade_req={:?}, downstream_was_upgraded={}, downstream_task_seen_upgraded={}, task_filter_seen_upgraded={}",
status.downstream,
status.upstream,
session.was_upgraded(),
session.downstream_task_seen_upgraded,
task_filter_seen_upgraded
),
)
.map_err(|e| e.into_in())
}
impl AsRef<HttpSession> for Session {
fn as_ref(&self) -> &HttpSession {
&self.downstream_session
}
}
impl AsMut<HttpSession> for Session {
fn as_mut(&mut self) -> &mut HttpSession {
&mut self.downstream_session
}
}
use std::ops::{Deref, DerefMut};
impl Deref for Session {
type Target = HttpSession;
fn deref(&self) -> &Self::Target {
&self.downstream_session
}
}
impl DerefMut for Session {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.downstream_session
}
}
static BAD_GATEWAY: Lazy<ResponseHeader> = Lazy::new(|| {
let mut resp = ResponseHeader::build(http::StatusCode::BAD_GATEWAY, Some(3)).unwrap();
resp.insert_header(header::SERVER, &SERVER_NAME[..])
.unwrap();
resp.insert_header(header::CONTENT_LENGTH, 0).unwrap();
resp.insert_header(header::CACHE_CONTROL, "private, no-store")
.unwrap();
resp
});
impl<SV, C> HttpProxy<SV, C>
where
C: custom::Connector,
{
async fn process_request(
self: &Arc<Self>,
mut session: Session,
mut ctx: <SV as ProxyHttp>::CTX,
) -> Option<ReusedHttpStream>
where
SV: ProxyHttp + Send + Sync + 'static,
<SV as ProxyHttp>::CTX: Send + Sync,
{
if let Err(e) = self
.inner
.early_request_filter(&mut session, &mut ctx)
.await
{
return self
.handle_error(session, &mut ctx, e, "Fail to early filter request:")
.await;
}
if self.inner.allow_spawning_subrequest(&session, &ctx) {
session.subrequest_spawner = Some(SubrequestSpawner::new(self.clone()));
}
let req = session.downstream_session.req_header_mut();
if let Err(e) = session
.downstream_modules_ctx
.request_header_filter(req)
.await
{
return self
.handle_error(
session,
&mut ctx,
e,
"Failed in downstream modules request filter:",
)
.await;
}
match self.inner.request_filter(&mut session, &mut ctx).await {
Ok(response_sent) => {
if response_sent {
self.inner.logging(&mut session, None, &mut ctx).await;
self.cleanup_sub_req(&mut session);
let mut persistent_settings = HttpPersistentSettings::for_session(&session);
if let Some(uc) = self.inner.persist_connection_context(&session, &ctx) {
persistent_settings.set_user_context(uc);
}
return session
.downstream_session
.finish()
.await
.ok()
.flatten()
.map(|s| ReusedHttpStream::from_reusable_stream(s, persistent_settings));
}
}
Err(e) => {
return self
.handle_error(session, &mut ctx, e, "Fail to filter request:")
.await;
}
}
if let Some((reuse, err)) = self.proxy_cache(&mut session, &mut ctx).await {
return self.finish(session, &mut ctx, reuse, err).await;
}
self.cleanup_sub_req(&mut session);
match self
.inner
.proxy_upstream_filter(&mut session, &mut ctx)
.await
{
Ok(proxy_to_upstream) => {
if !proxy_to_upstream {
if session.cache.enabled() {
session.cache.disable(NoCacheReason::DeclinedToUpstream);
}
if session.response_written().is_none() {
match session.write_response_header_ref(&BAD_GATEWAY, true).await {
Ok(()) => {}
Err(e) => {
return self
.handle_error(
session,
&mut ctx,
e,
"Error responding with Bad Gateway:",
)
.await;
}
}
}
return self.finish(session, &mut ctx, true, None).await;
}
}
Err(e) => {
if session.cache.enabled() {
session.cache.disable(NoCacheReason::InternalError);
}
return self
.handle_error(
session,
&mut ctx,
e,
"Error deciding if we should proxy to upstream:",
)
.await;
}
}
let mut retries: usize = 0;
let mut server_reuse = false;
let mut proxy_error: Option<Box<Error>> = None;
while retries < self.max_retries {
retries += 1;
let (reuse, e) = self.proxy_to_upstream(&mut session, &mut ctx).await;
server_reuse = reuse;
match e {
Some(error) => {
let retry = error.retry();
if retry
&& !self.inner.suppress_proxy_warn_log(
&session,
&ctx,
&error,
ProxyWarnLogContext::UpstreamRetry,
)
{
warn!(
"Fail to proxy: {}, tries: {}, retry: {}, {}",
error,
retries,
retry,
self.inner.request_summary(&session, &ctx)
);
}
proxy_error = Some(error);
if !retry {
break;
}
}
None => {
proxy_error = None;
break;
}
};
}
#[allow(clippy::unnecessary_unwrap)]
let serve_stale_result = if proxy_error.is_some() && session.cache.can_serve_stale_error() {
self.handle_stale_if_error(&mut session, &mut ctx, proxy_error.as_ref().unwrap())
.await
} else {
None
};
let final_error = if let Some((reuse, stale_cache_error)) = serve_stale_result {
server_reuse = server_reuse && reuse;
stale_cache_error
} else {
proxy_error
};
if let Some(e) = final_error.as_ref() {
if session.cache.enabled() {
let reason = if *e.esource() == ErrorSource::Upstream {
NoCacheReason::UpstreamError
} else {
NoCacheReason::InternalError
};
session.cache.disable(reason);
}
let res = self.inner.fail_to_proxy(&mut session, e, &mut ctx).await;
if !self.inner.suppress_error_log(&session, &ctx, e) {
error!(
"Fail to proxy: {}, status: {}, tries: {}, retry: {}, {}",
e,
res.error_code,
retries,
false, self.inner.request_summary(&session, &ctx),
);
}
}
self.finish(session, &mut ctx, server_reuse, final_error)
.await
}
async fn handle_error(
&self,
mut session: Session,
ctx: &mut <SV as ProxyHttp>::CTX,
e: Box<Error>,
context: &str,
) -> Option<ReusedHttpStream>
where
SV: ProxyHttp + Send + Sync + 'static,
<SV as ProxyHttp>::CTX: Send + Sync,
{
let res = self.inner.fail_to_proxy(&mut session, &e, ctx).await;
if !self.inner.suppress_error_log(&session, ctx, &e) {
error!(
"{context} {}, status: {}, {}",
e,
res.error_code,
self.inner.request_summary(&session, ctx)
);
}
self.inner.logging(&mut session, Some(&e), ctx).await;
self.cleanup_sub_req(&mut session);
session.downstream_session.on_proxy_failure(e);
if res.can_reuse_downstream {
let mut persistent_settings = HttpPersistentSettings::for_session(&session);
if let Some(uc) = self.inner.persist_connection_context(&session, ctx) {
persistent_settings.set_user_context(uc);
}
session
.downstream_session
.finish()
.await
.ok()
.flatten()
.map(|s| ReusedHttpStream::from_reusable_stream(s, persistent_settings))
} else {
None
}
}
}
#[async_trait]
pub trait Subrequest {
async fn process_subrequest(
self: Arc<Self>,
session: Box<HttpSession>,
sub_req_ctx: Box<SubrequestCtx>,
);
}
#[async_trait]
impl<SV, C> Subrequest for HttpProxy<SV, C>
where
SV: ProxyHttp + Send + Sync + 'static,
<SV as ProxyHttp>::CTX: Send + Sync,
C: custom::Connector,
{
async fn process_subrequest(
self: Arc<Self>,
session: Box<HttpSession>,
sub_req_ctx: Box<SubrequestCtx>,
) {
debug!("starting subrequest");
let mut session = match self.handle_new_request(session).await {
Some(downstream_session) => Session::new(
downstream_session,
&self.downstream_modules,
#[cfg(feature = "upstream_modules")]
&self.upstream_modules,
self.shutdown_flag.clone(),
),
None => return, };
session.set_keepalive(None);
session.subrequest_ctx.replace(sub_req_ctx);
trace!("processing subrequest");
let ctx = self.inner.new_ctx();
self.process_request(session, ctx).await;
trace!("subrequest done");
}
}
pub struct SubrequestSpawner {
app: Arc<dyn Subrequest + Send + Sync>,
}
pub struct PreparedSubrequest {
app: Arc<dyn Subrequest + Send + Sync>,
session: Box<HttpSession>,
sub_req_ctx: Box<SubrequestCtx>,
}
impl PreparedSubrequest {
pub async fn run(self) {
self.app
.process_subrequest(self.session, self.sub_req_ctx)
.await
}
pub fn session(&self) -> &HttpSession {
self.session.as_ref()
}
pub fn session_mut(&mut self) -> &mut HttpSession {
self.session.deref_mut()
}
}
impl SubrequestSpawner {
pub fn new(app: Arc<dyn Subrequest + Send + Sync>) -> SubrequestSpawner {
SubrequestSpawner { app }
}
pub fn spawn_background_subrequest(
&self,
session: &HttpSession,
ctx: SubrequestCtx,
) -> tokio::task::JoinHandle<()> {
let new_app = self.app.clone(); let (mut session, handle) = subrequest::create_session(session);
if ctx.body_mode() == BodyMode::NoBody {
session
.as_subrequest_mut()
.expect("created subrequest session")
.clear_request_body_headers();
}
let sub_req_ctx = Box::new(ctx);
handle.drain_tasks();
tokio::spawn(async move {
new_app
.process_subrequest(Box::new(session), sub_req_ctx)
.await;
})
}
pub fn create_subrequest(
&self,
session: &HttpSession,
ctx: SubrequestCtx,
) -> (PreparedSubrequest, SubrequestHandle) {
let new_app = self.app.clone(); let (mut session, handle) = subrequest::create_session(session);
if ctx.body_mode() == BodyMode::NoBody {
session
.as_subrequest_mut()
.expect("created subrequest session")
.clear_request_body_headers();
}
let sub_req_ctx = Box::new(ctx);
(
PreparedSubrequest {
app: new_app,
session: Box::new(session),
sub_req_ctx,
},
handle,
)
}
}
#[async_trait]
impl<SV, C> HttpServerApp for HttpProxy<SV, C>
where
SV: ProxyHttp + Send + Sync + 'static,
<SV as ProxyHttp>::CTX: Send + Sync,
C: custom::Connector,
{
async fn process_new_http(
self: &Arc<Self>,
mut session: HttpSession,
shutdown: &ShutdownWatch,
) -> Option<ReusedHttpStream> {
let prev_user_ctx = session.take_connection_user_context();
let session = Box::new(session);
let mut session = match self.handle_new_request(session).await {
Some(downstream_session) => Session::new(
downstream_session,
&self.downstream_modules,
#[cfg(feature = "upstream_modules")]
&self.upstream_modules,
self.shutdown_flag.clone(),
),
None => return None, };
if *shutdown.borrow() {
session.set_keepalive(None);
}
let mut ctx = self.inner.new_ctx();
if let Some(prev_ctx) = prev_user_ctx {
self.inner
.on_connection_reuse(&mut session, &mut ctx, prev_ctx);
}
self.process_request(session, ctx).await
}
async fn http_cleanup(&self) {
self.shutdown_flag.store(true, Ordering::Release);
self.shutdown.notify_waiters();
}
fn server_options(&self) -> Option<&HttpServerOptions> {
self.server_options.as_ref()
}
fn h2_options(&self) -> Option<H2Options> {
self.h2_options.clone()
}
async fn process_custom_session(
self: Arc<Self>,
stream: Stream,
shutdown: &ShutdownWatch,
) -> Option<Stream> {
let app = self.clone();
let Some(process_custom_session) = app.process_custom_session.as_ref() else {
warn!("custom was called on an empty on_custom");
return None;
};
process_custom_session(self.clone(), stream, shutdown).await
}
}
use pingora_core::services::listening::{RuntimeOptsOverride, Service};
pub fn http_proxy<SV>(conf: &Arc<ServerConf>, inner: SV) -> HttpProxy<SV>
where
SV: ProxyHttp,
{
let mut proxy = HttpProxy::new(inner, conf.clone());
proxy.handle_init_modules();
proxy
}
pub fn http_proxy_service<SV>(conf: &Arc<ServerConf>, inner: SV) -> Service<HttpProxy<SV, ()>>
where
SV: ProxyHttp,
{
http_proxy_service_with_name(conf, inner, "Pingora HTTP Proxy Service")
}
pub fn http_proxy_service_with_name<SV>(
conf: &Arc<ServerConf>,
inner: SV,
name: &str,
) -> Service<HttpProxy<SV, ()>>
where
SV: ProxyHttp,
{
let mut proxy = HttpProxy::new(inner, conf.clone());
proxy.handle_init_modules();
Service::new(name.to_string(), proxy)
}
pub fn http_proxy_service_with_name_custom<SV, C>(
conf: &Arc<ServerConf>,
inner: SV,
name: &str,
connector: C,
on_custom: ProcessCustomSession<SV, C>,
) -> Service<HttpProxy<SV, C>>
where
SV: ProxyHttp + Send + Sync + 'static,
SV::CTX: Send + Sync + 'static,
C: custom::Connector,
{
let mut proxy =
HttpProxy::new_custom(inner, conf.clone(), connector, Some(on_custom), None, None);
proxy.handle_init_modules();
Service::new(name.to_string(), proxy)
}
pub struct ProxyServiceBuilder<SV, C>
where
SV: ProxyHttp + Send + Sync + 'static,
SV::CTX: Send + Sync + 'static,
C: custom::Connector,
{
conf: Arc<ServerConf>,
inner: SV,
name: String,
connector: C,
custom: Option<ProcessCustomSession<SV, C>>,
server_options: Option<HttpServerOptions>,
client_options: Option<ConnectorOptions>,
runtime_opts_override: Option<RuntimeOptsOverride>,
}
impl<SV> ProxyServiceBuilder<SV, ()>
where
SV: ProxyHttp + Send + Sync + 'static,
SV::CTX: Send + Sync + 'static,
{
pub fn new(conf: &Arc<ServerConf>, inner: SV) -> Self {
ProxyServiceBuilder {
conf: conf.clone(),
inner,
name: "Pingora HTTP Proxy Service".into(),
connector: (),
custom: None,
server_options: None,
client_options: None,
runtime_opts_override: None,
}
}
}
impl<SV, C> ProxyServiceBuilder<SV, C>
where
SV: ProxyHttp + Send + Sync + 'static,
SV::CTX: Send + Sync + 'static,
C: custom::Connector,
{
pub fn name(mut self, name: impl AsRef<str>) -> Self {
self.name = name.as_ref().to_owned();
self
}
pub fn custom<C2: custom::Connector>(
self,
connector: C2,
on_custom: ProcessCustomSession<SV, C2>,
) -> ProxyServiceBuilder<SV, C2> {
let Self {
conf,
inner,
name,
server_options,
client_options,
runtime_opts_override,
..
} = self;
ProxyServiceBuilder {
conf,
inner,
name,
connector,
custom: Some(on_custom),
server_options,
client_options,
runtime_opts_override,
}
}
pub fn client_options(mut self, options: ConnectorOptions) -> Self {
self.client_options = Some(options);
self
}
pub fn server_options(mut self, options: HttpServerOptions) -> Self {
self.server_options = Some(options);
self
}
pub fn runtime_opts_override<F>(mut self, override_fn: F) -> Self
where
F: Fn(&RuntimeOpts) -> Option<RuntimeOpts> + Send + Sync + 'static,
{
self.runtime_opts_override = Some(Arc::new(override_fn));
self
}
pub fn build(self) -> Service<HttpProxy<SV, C>> {
let Self {
conf,
inner,
name,
connector,
custom,
server_options,
client_options,
runtime_opts_override,
} = self;
let mut proxy = HttpProxy::new_custom(
inner,
conf,
connector,
custom,
server_options,
client_options,
);
proxy.handle_init_modules();
let mut service = Service::new(name, proxy);
if let Some(runtime_opts_override) = runtime_opts_override {
service.set_runtime_opts_override(runtime_opts_override);
}
service
}
}
#[cfg(test)]
mod tests {
use super::*;
use pingora_core::modules::http::{HttpModule, HttpModuleBuilder};
use pingora_core::protocols::l4::stream::Stream as L4Stream;
use pingora_core::protocols::l4::virt::{VirtualSockOpt, VirtualSocket, VirtualSocketStream};
use pingora_error::RetryType;
use std::pin::Pin;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Mutex;
use std::task::{Context, Poll};
use tokio::io::{AsyncRead, AsyncWrite, ReadBuf};
#[derive(Debug)]
struct StaticVirtualSocket {
read_buf: Vec<u8>,
read_pos: usize,
write_buf: Arc<Mutex<Vec<u8>>>,
}
impl StaticVirtualSocket {
fn new(read_buf: &[u8], write_buf: Arc<Mutex<Vec<u8>>>) -> Self {
Self {
read_buf: read_buf.to_vec(),
read_pos: 0,
write_buf,
}
}
}
impl AsyncRead for StaticVirtualSocket {
fn poll_read(
mut self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
let remaining = self.read_buf.len() - self.read_pos;
let to_read = remaining.min(buf.remaining());
if to_read > 0 {
buf.put_slice(&self.read_buf[self.read_pos..self.read_pos + to_read]);
self.read_pos += to_read;
}
Poll::Ready(Ok(()))
}
}
impl AsyncWrite for StaticVirtualSocket {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
self.write_buf.lock().unwrap().extend_from_slice(buf);
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Poll::Ready(Ok(()))
}
}
impl VirtualSocket for StaticVirtualSocket {
fn set_socket_option(&self, _opt: VirtualSockOpt) -> std::io::Result<()> {
Ok(())
}
}
async fn new_request_session(request: &[u8], written: Arc<Mutex<Vec<u8>>>) -> Session {
let socket = StaticVirtualSocket::new(request, written);
let stream = L4Stream::from(VirtualSocketStream::new(Box::new(socket)));
let mut session = Session::new_h1(Box::new(stream));
session.read_request().await.unwrap();
session
}
async fn new_upgrade_request_session(written: Arc<Mutex<Vec<u8>>>) -> Session {
new_request_session(
b"GET / HTTP/1.1\r\nHost: example.com\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n",
written,
)
.await
}
struct DefaultRetryProxy;
#[async_trait]
impl ProxyHttp for DefaultRetryProxy {
type CTX = ();
fn new_ctx(&self) -> Self::CTX {}
async fn upstream_peer(
&self,
_session: &mut Session,
_ctx: &mut Self::CTX,
) -> Result<Box<HttpPeer>> {
unreachable!()
}
}
fn default_policy_would_retry_for_session(
session: &mut Session,
retry: RetryType,
client_reused: bool,
) -> bool {
let mut error = Error::new_up(ReadError);
error.retry = retry;
DefaultRetryProxy
.error_while_proxy(
&HttpPeer::new("127.0.0.1:80", false, "".to_string()),
session,
error,
&mut (),
client_reused,
)
.retry()
}
async fn default_policy_would_retry(
request: &[u8],
retry: RetryType,
client_reused: bool,
) -> bool {
let mut session = new_request_session(request, Arc::new(Mutex::new(Vec::new()))).await;
default_policy_would_retry_for_session(&mut session, retry, client_reused)
}
async fn buffered_put_session(body_len: usize) -> Session {
let mut request =
format!("PUT / HTTP/1.1\r\nHost: example.com\r\nContent-Length: {body_len}\r\n\r\n")
.into_bytes();
request.resize(request.len() + body_len, b'a');
let mut session = new_request_session(&request, Arc::new(Mutex::new(Vec::new()))).await;
session.enable_retry_buffering();
while session.read_request_body().await.unwrap().is_some() {}
session
}
#[tokio::test]
async fn default_retry_policy_requires_an_idempotent_method() {
let decided_retry = RetryType::Decided(true);
assert!(
default_policy_would_retry(
b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n",
decided_retry,
false,
)
.await
);
assert!(
default_policy_would_retry(
b"PUT / HTTP/1.1\r\nHost: example.com\r\nContent-Length: 0\r\n\r\n",
decided_retry,
false,
)
.await
);
assert!(
!default_policy_would_retry(
b"POST / HTTP/1.1\r\nHost: example.com\r\nContent-Length: 0\r\n\r\n",
decided_retry,
false,
)
.await
);
assert!(
!default_policy_would_retry(
b"PATCH / HTTP/1.1\r\nHost: example.com\r\nContent-Length: 0\r\n\r\n",
decided_retry,
false,
)
.await
);
}
#[tokio::test]
async fn default_retry_policy_resolves_reused_only() {
let request = b"GET / HTTP/1.1\r\nHost: example.com\r\n\r\n";
assert!(default_policy_would_retry(request, RetryType::ReusedOnly, true).await);
assert!(!default_policy_would_retry(request, RetryType::ReusedOnly, false).await);
}
#[tokio::test]
async fn default_retry_policy_requires_an_untruncated_body_buffer() {
let mut complete = buffered_put_session(64 * 1024).await;
assert!(!complete.retry_buffer_truncated());
assert!(default_policy_would_retry_for_session(
&mut complete,
RetryType::Decided(true),
false,
));
let mut truncated = buffered_put_session(64 * 1024 + 1).await;
assert!(truncated.retry_buffer_truncated());
assert!(!default_policy_would_retry_for_session(
&mut truncated,
RetryType::Decided(true),
false,
));
assert!(!default_policy_would_retry_for_session(
&mut truncated,
RetryType::ReusedOnly,
true,
));
}
fn upgrade_response_header() -> ResponseHeader {
let mut header =
ResponseHeader::build(http::StatusCode::SWITCHING_PROTOCOLS, Some(2)).unwrap();
header
.insert_header(http::header::UPGRADE, "websocket")
.unwrap();
header
.insert_header(http::header::CONNECTION, "Upgrade")
.unwrap();
header
}
struct SwitchTo101Module;
#[async_trait]
impl HttpModule for SwitchTo101Module {
async fn response_header_filter(
&mut self,
resp: &mut ResponseHeader,
_end_of_stream: bool,
) -> Result<()> {
resp.set_status(http::StatusCode::SWITCHING_PROTOCOLS)?;
resp.set_version(Version::HTTP_11);
Ok(())
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
self
}
}
struct SwitchTo101ModuleBuilder;
impl HttpModuleBuilder for SwitchTo101ModuleBuilder {
fn init(&self) -> pingora_core::modules::http::Module {
Box::new(SwitchTo101Module)
}
}
struct DoneBytesModule {
called: Arc<AtomicBool>,
}
#[async_trait]
impl HttpModule for DoneBytesModule {
fn response_done_filter(&mut self) -> Result<Option<Bytes>> {
self.called.store(true, Ordering::Release);
Ok(Some(Bytes::from_static(b"hello")))
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
self
}
}
struct DoneBytesModuleBuilder {
called: Arc<AtomicBool>,
}
impl HttpModuleBuilder for DoneBytesModuleBuilder {
fn init(&self) -> pingora_core::modules::http::Module {
Box::new(DoneBytesModule {
called: self.called.clone(),
})
}
}
struct DoneEmptyModule {
called: Arc<AtomicBool>,
}
impl HttpModule for DoneEmptyModule {
fn response_done_filter(&mut self) -> Result<Option<Bytes>> {
self.called.store(true, Ordering::Release);
Ok(None)
}
fn as_any(&self) -> &dyn std::any::Any {
self
}
fn as_any_mut(&mut self) -> &mut dyn std::any::Any {
self
}
}
struct DoneEmptyModuleBuilder {
called: Arc<AtomicBool>,
}
impl HttpModuleBuilder for DoneEmptyModuleBuilder {
fn init(&self) -> pingora_core::modules::http::Module {
Box::new(DoneEmptyModule {
called: self.called.clone(),
})
}
}
fn assert_raw_upgrade_payload(written: &[u8]) {
assert!(
written.starts_with(b"HTTP/1.1 101 Switching Protocols\r\n"),
"unexpected response: {:?}",
String::from_utf8_lossy(written)
);
assert!(
written.ends_with(b"\r\n\r\nhello"),
"upgrade payload should be written as raw tunneled bytes: {:?}",
String::from_utf8_lossy(written)
);
assert!(
!written
.windows(b"\r\n5\r\nhello".len())
.any(|w| w == b"\r\n5\r\nhello"),
"upgrade payload must not be chunk framed: {:?}",
String::from_utf8_lossy(written)
);
}
#[tokio::test]
async fn write_response_tasks_rejects_body_after_101() {
let written = Arc::new(Mutex::new(Vec::new()));
let mut session = new_upgrade_request_session(written.clone()).await;
let err = session
.write_response_tasks(vec![
HttpTask::Header(Box::new(upgrade_response_header()), false),
HttpTask::Body(Some(Bytes::from_static(b"hello")), true),
])
.await
.unwrap_err();
assert_eq!(err.etype(), &InvalidHTTPHeader);
assert_eq!(err.esource(), &ErrorSource::Internal);
assert!(written.lock().unwrap().is_empty());
}
#[tokio::test]
async fn write_response_tasks_allows_upgraded_body_after_101() {
let written = Arc::new(Mutex::new(Vec::new()));
let mut session = new_upgrade_request_session(written.clone()).await;
let response_done = session
.write_response_tasks(vec![
HttpTask::Header(Box::new(upgrade_response_header()), false),
HttpTask::UpgradedBody(Some(Bytes::from_static(b"hello")), true),
])
.await
.unwrap();
assert!(response_done);
let written = written.lock().unwrap().clone();
assert_raw_upgrade_payload(&written);
}
#[tokio::test]
async fn write_response_tasks_rejects_upgraded_body_before_101() {
let written = Arc::new(Mutex::new(Vec::new()));
let mut session = new_upgrade_request_session(written.clone()).await;
session.set_upstream_h1_upgrade_request_status(true);
let err = session
.write_response_tasks(vec![HttpTask::UpgradedBody(
Some(Bytes::from_static(b"hello")),
true,
)])
.await
.unwrap_err();
assert_eq!(err.etype(), &InvalidHTTPHeader);
assert_eq!(err.esource(), &ErrorSource::Internal);
assert!(written.lock().unwrap().is_empty());
}
#[tokio::test]
async fn write_response_tasks_rejects_trailer_after_101() {
let written = Arc::new(Mutex::new(Vec::new()));
let mut session = new_upgrade_request_session(written.clone()).await;
let err = session
.write_response_tasks(vec![
HttpTask::Header(Box::new(upgrade_response_header()), false),
HttpTask::Trailer(Some(Box::new(http::HeaderMap::new()))),
])
.await
.unwrap_err();
assert_eq!(err.etype(), &InvalidHTTPHeader);
assert_eq!(err.esource(), &ErrorSource::Internal);
assert!(written.lock().unwrap().is_empty());
}
#[tokio::test]
async fn write_response_tasks_runs_done_filter_after_101_as_upgraded_body() {
let written = Arc::new(Mutex::new(Vec::new()));
let called = Arc::new(AtomicBool::new(false));
let socket = StaticVirtualSocket::new(
b"GET / HTTP/1.1\r\nHost: example.com\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n",
written.clone(),
);
let stream = L4Stream::from(VirtualSocketStream::new(Box::new(socket)));
let mut modules = HttpModules::new();
modules.add_module(Box::new(DoneBytesModuleBuilder {
called: called.clone(),
}));
let mut session = Session::new_h1_with_modules(Box::new(stream), &modules);
session.read_request().await.unwrap();
let response_done = session
.write_response_tasks(vec![
HttpTask::Header(Box::new(upgrade_response_header()), false),
HttpTask::Done,
])
.await
.unwrap();
assert!(response_done);
assert!(called.load(Ordering::Acquire));
let written = written.lock().unwrap().clone();
assert_raw_upgrade_payload(&written);
}
#[tokio::test]
async fn write_response_tasks_allows_empty_done_after_101() {
let written = Arc::new(Mutex::new(Vec::new()));
let called = Arc::new(AtomicBool::new(false));
let socket = StaticVirtualSocket::new(
b"GET / HTTP/1.1\r\nHost: example.com\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n",
written.clone(),
);
let stream = L4Stream::from(VirtualSocketStream::new(Box::new(socket)));
let mut modules = HttpModules::new();
modules.add_module(Box::new(DoneEmptyModuleBuilder {
called: called.clone(),
}));
let mut session = Session::new_h1_with_modules(Box::new(stream), &modules);
session.read_request().await.unwrap();
let response_done = session
.write_response_tasks(vec![
HttpTask::Header(Box::new(upgrade_response_header()), false),
HttpTask::Done,
])
.await
.unwrap();
assert!(response_done);
assert!(called.load(Ordering::Acquire));
let written = written.lock().unwrap().clone();
assert!(
written.starts_with(b"HTTP/1.1 101 Switching Protocols\r\n"),
"unexpected response: {:?}",
String::from_utf8_lossy(&written)
);
assert!(
written.ends_with(b"\r\n\r\n"),
"empty Done filter should only finish the upgraded response: {:?}",
String::from_utf8_lossy(&written)
);
}
#[tokio::test]
async fn write_response_tasks_rejects_module_created_101_with_upgrade_mismatch() {
let written = Arc::new(Mutex::new(Vec::new()));
let socket = StaticVirtualSocket::new(
b"GET / HTTP/1.1\r\nHost: example.com\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n",
written.clone(),
);
let stream = L4Stream::from(VirtualSocketStream::new(Box::new(socket)));
let mut modules = HttpModules::new();
modules.add_module(Box::new(SwitchTo101ModuleBuilder));
let mut session = Session::new_h1_with_modules(Box::new(stream), &modules);
session.read_request().await.unwrap();
session.h1_upgrade_request_status = H1UpgradeRequestStatus {
upstream: Some(false),
};
let err = session
.write_response_tasks(vec![
HttpTask::Header(
Box::new(ResponseHeader::build(200, Some(0)).unwrap()),
false,
),
HttpTask::Body(Some(Bytes::from_static(b"hello")), true),
])
.await
.unwrap_err();
assert_eq!(err.etype(), &InvalidHTTPHeader);
assert_eq!(err.esource(), &ErrorSource::Internal);
assert!(written.lock().unwrap().is_empty());
}
#[tokio::test]
async fn write_response_tasks_rejects_module_created_101_before_body() {
let written = Arc::new(Mutex::new(Vec::new()));
let socket = StaticVirtualSocket::new(
b"GET / HTTP/1.1\r\nHost: example.com\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n",
written.clone(),
);
let stream = L4Stream::from(VirtualSocketStream::new(Box::new(socket)));
let mut modules = HttpModules::new();
modules.add_module(Box::new(SwitchTo101ModuleBuilder));
let mut session = Session::new_h1_with_modules(Box::new(stream), &modules);
session.read_request().await.unwrap();
let err = session
.write_response_tasks(vec![
HttpTask::Header(
Box::new(ResponseHeader::build(200, Some(0)).unwrap()),
false,
),
HttpTask::Body(Some(Bytes::from_static(b"hello")), true),
])
.await
.unwrap_err();
assert_eq!(err.etype(), &InvalidHTTPHeader);
assert_eq!(err.esource(), &ErrorSource::Internal);
assert!(written.lock().unwrap().is_empty());
}
#[tokio::test]
async fn send_downstream_proxy_task_rejects_body_after_101() {
let written = Arc::new(Mutex::new(Vec::new()));
let mut session = new_upgrade_request_session(written.clone()).await;
session.set_proxy_tasks_enabled(true);
session
.send_downstream_proxy_task(HttpTask::Header(
Box::new(upgrade_response_header()),
false,
))
.await
.unwrap();
let err = session
.send_downstream_proxy_task(HttpTask::Body(Some(Bytes::from_static(b"hello")), true))
.await
.unwrap_err();
assert_eq!(err.etype(), &InvalidHTTPHeader);
assert_eq!(err.esource(), &ErrorSource::Internal);
assert!(written.lock().unwrap().is_empty());
}
#[tokio::test]
async fn send_downstream_proxy_task_allows_upgraded_body_after_101() {
let written = Arc::new(Mutex::new(Vec::new()));
let mut session = new_upgrade_request_session(written.clone()).await;
session.set_proxy_tasks_enabled(true);
session
.send_downstream_proxy_task(HttpTask::Header(
Box::new(upgrade_response_header()),
false,
))
.await
.unwrap();
session
.send_downstream_proxy_task(HttpTask::UpgradedBody(
Some(Bytes::from_static(b"hello")),
true,
))
.await
.unwrap();
let response_done = session.write_downstream_proxy_tasks().await.unwrap();
assert!(response_done);
let written = written.lock().unwrap().clone();
assert_raw_upgrade_payload(&written);
}
#[tokio::test]
async fn send_downstream_proxy_task_rejects_upgraded_body_before_101() {
let written = Arc::new(Mutex::new(Vec::new()));
let mut session = new_upgrade_request_session(written.clone()).await;
session.set_upstream_h1_upgrade_request_status(true);
session.set_proxy_tasks_enabled(true);
let err = session
.send_downstream_proxy_task(HttpTask::UpgradedBody(
Some(Bytes::from_static(b"hello")),
true,
))
.await
.unwrap_err();
assert_eq!(err.etype(), &InvalidHTTPHeader);
assert_eq!(err.esource(), &ErrorSource::Internal);
assert!(!session.has_pending_downstream_tasks());
assert!(written.lock().unwrap().is_empty());
}
#[tokio::test]
async fn send_downstream_proxy_task_runs_done_filter_after_101_as_upgraded_body() {
let written = Arc::new(Mutex::new(Vec::new()));
let called = Arc::new(AtomicBool::new(false));
let socket = StaticVirtualSocket::new(
b"GET / HTTP/1.1\r\nHost: example.com\r\nUpgrade: websocket\r\nConnection: Upgrade\r\n\r\n",
written.clone(),
);
let stream = L4Stream::from(VirtualSocketStream::new(Box::new(socket)));
let mut modules = HttpModules::new();
modules.add_module(Box::new(DoneBytesModuleBuilder {
called: called.clone(),
}));
let mut session = Session::new_h1_with_modules(Box::new(stream), &modules);
session.read_request().await.unwrap();
session.set_proxy_tasks_enabled(true);
session
.send_downstream_proxy_task(HttpTask::Header(
Box::new(upgrade_response_header()),
false,
))
.await
.unwrap();
session
.send_downstream_proxy_task(HttpTask::Done)
.await
.unwrap();
let response_done = session.write_downstream_proxy_tasks().await.unwrap();
assert!(response_done);
assert!(called.load(Ordering::Acquire));
let written = written.lock().unwrap().clone();
assert_raw_upgrade_payload(&written);
}
#[derive(Debug)]
struct PendingVirtualSocket;
impl AsyncRead for PendingVirtualSocket {
fn poll_read(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
_buf: &mut ReadBuf<'_>,
) -> Poll<std::io::Result<()>> {
Poll::Pending
}
}
impl AsyncWrite for PendingVirtualSocket {
fn poll_write(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
buf: &[u8],
) -> Poll<std::io::Result<usize>> {
Poll::Ready(Ok(buf.len()))
}
fn poll_flush(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Poll::Ready(Ok(()))
}
fn poll_shutdown(self: Pin<&mut Self>, _cx: &mut Context<'_>) -> Poll<std::io::Result<()>> {
Poll::Ready(Ok(()))
}
}
impl VirtualSocket for PendingVirtualSocket {
fn set_socket_option(&self, _opt: VirtualSockOpt) -> std::io::Result<()> {
Ok(())
}
}
struct NoopProxy;
#[async_trait]
impl ProxyHttp for NoopProxy {
type CTX = ();
fn new_ctx(&self) -> Self::CTX {}
async fn upstream_peer(
&self,
_session: &mut Session,
_ctx: &mut Self::CTX,
) -> Result<Box<HttpPeer>> {
Err(Error::new(InternalError))
}
}
fn pending_session() -> Box<HttpSession> {
let stream = L4Stream::from(VirtualSocketStream::new(Box::new(PendingVirtualSocket)));
Box::new(HttpSession::new_http1(Box::new(stream)))
}
#[tokio::test(flavor = "multi_thread", worker_threads = 4)]
async fn shutdown_wakes_parked_read_requests() {
let conf = ServerConf {
threads: 4,
..ServerConf::default()
};
let proxy = Arc::new(HttpProxy::new(NoopProxy, Arc::new(conf)));
let handles: Vec<_> = (0..8)
.map(|_| {
let proxy = proxy.clone();
tokio::spawn(async move { proxy.handle_new_request(pending_session()).await })
})
.collect();
time::sleep(Duration::from_millis(50)).await;
proxy.http_cleanup().await;
for handle in handles {
let session = time::timeout(Duration::from_secs(5), handle)
.await
.expect("shutdown did not wake the parked read")
.unwrap();
assert!(session.is_none());
}
}
#[tokio::test]
async fn shutdown_before_read_request_parks_returns_immediately() {
let proxy = Arc::new(HttpProxy::new(NoopProxy, Arc::new(ServerConf::default())));
proxy.http_cleanup().await;
let session = time::timeout(
Duration::from_secs(5),
proxy.handle_new_request(pending_session()),
)
.await
.expect("read_request parked after shutdown");
assert!(session.is_none());
}
}