use std::{
collections::{HashMap, VecDeque},
time::Duration,
};
use async_trait::async_trait;
use bytes::Bytes;
use praxis_core::subrequest::SubResponseBody;
use tracing::warn;
use super::{continuation::FilteredSubrequestContinuation, transport::termination_cause};
use crate::{
FilterError, StreamTermination, StreamTerminationCause, actions::StreamingResponseBody,
context::PendingStreamChunks, extensions::RequestExtensions,
};
pub(crate) struct FilteredStreamingBody {
upstream: Option<Box<SubResponseBody>>,
continuation: FilteredSubrequestContinuation,
finished: bool,
deferred_completion_output: Option<Bytes>,
pending_chunks: VecDeque<Bytes>,
}
impl FilteredStreamingBody {
pub(crate) fn new(upstream: Box<SubResponseBody>, continuation: FilteredSubrequestContinuation) -> Self {
Self {
upstream: Some(upstream),
continuation,
finished: false,
deferred_completion_output: None,
pending_chunks: VecDeque::new(),
}
}
pub(crate) fn into_continuation(self) -> FilteredSubrequestContinuation {
self.continuation
}
pub(crate) fn into_finished_parts(self) -> (FilteredSubrequestContinuation, Option<Bytes>) {
(self.continuation, self.deferred_completion_output)
}
fn exchange_extensions(&mut self, extensions: &mut RequestExtensions) {
std::mem::swap(&mut self.continuation.extensions, extensions);
}
#[expect(clippy::too_many_lines, reason = "context reconstruction requires many fields")]
fn run_step_body_filters(&mut self, body: &mut Option<Bytes>, end_of_stream: bool) -> Result<(), FilterError> {
let remaining_read_timeout;
let stream_deadline;
let result = {
let cont = &mut self.continuation;
let mut ctx = crate::filter::HttpFilterContext {
grpc_completion: None,
buffered_request_body: None,
body_done_indices: std::mem::take(&mut cont.body_done_indices),
branch_iterations: HashMap::new(),
client_addr: cont.client_addr,
cluster: None,
current_filter_id: None,
downstream_tls: cont.downstream_tls,
extensions: std::mem::take(&mut cont.extensions),
executed_branch_filters: Vec::new(),
executed_filter_indices: std::mem::take(&mut cont.executed_filter_indices),
extra_request_headers: Vec::new(),
filter_metadata: std::mem::take(&mut cont.filter_metadata),
filter_results: std::mem::take(&mut cont.filter_results),
filter_state: std::mem::take(&mut cont.filter_state),
health_registry: cont.pipeline.health_registry(),
id_generator: cont.pipeline.id_generator(),
kv_stores: cont.pipeline.kv_stores(),
metrics_route: None,
peer_identity: cont.peer_identity.clone(),
prior_pre_read_mutations: Vec::new(),
pre_read_mutations: Vec::new(),
request: &cont.request_snapshot,
request_body_bytes: 0,
request_body_mode: cont.pipeline.body_capabilities().request_body_mode,
request_headers_to_remove: Vec::new(),
request_headers_to_set: Vec::new(),
request_start: cont.request_start,
response_body_bytes: cont.response_body_bytes,
response_body_mode: cont.response_body_mode,
response_header: None,
response_headers_modified: false,
upstream_reached: false,
rewritten_path: None,
selected_endpoint_index: None,
attempted_endpoints: Vec::new(),
retry_policy: None,
route_retry_policy: None,
cluster_retry_state: None,
cluster_retry_state_released: false,
endpoint_reselector: None,
pinned_endpoint_address: None,
session_stores: cont.pipeline.session_stores(),
structured_metadata: std::mem::take(&mut cont.structured_metadata),
subrequest_client: cont.pipeline.subrequest_client(),
subrequest_response_mode: crate::context::SubRequestResponseMode::Streaming,
time_source: cont.pipeline.time_source(),
upstream: None,
};
let result = cont.pipeline.execute_http_response_body_with_response_header(
&mut ctx,
body,
end_of_stream,
Some(&cont.response_snapshot),
);
remaining_read_timeout = leftover_stream_read_timeout(&mut ctx);
stream_deadline = leftover_stream_deadline(&mut ctx);
cont.body_done_indices = ctx.body_done_indices;
cont.executed_filter_indices = ctx.executed_filter_indices;
cont.extensions = ctx.extensions;
cont.filter_metadata = ctx.filter_metadata;
cont.filter_results = ctx.filter_results;
cont.filter_state = ctx.filter_state;
cont.response_body_bytes = ctx.response_body_bytes;
cont.structured_metadata = ctx.structured_metadata;
result
};
if let crate::actions::FilterAction::Reject(_) = result? {
return Err("filtered_subrequest: step body filter rejected during stream"
.to_owned()
.into());
}
apply_leftover_read_timeout(&mut self.upstream, remaining_read_timeout);
apply_leftover_stream_deadline(&mut self.upstream, stream_deadline);
Ok(())
}
fn complete_step(&mut self) -> Result<Option<Bytes>, FilterError> {
if self.continuation.completed {
return Ok(None);
}
self.continuation.completed = true;
let mut body: Option<Bytes> = None;
self.run_step_body_filters(&mut body, true)?;
Ok(body)
}
fn handle_upstream_chunk(&mut self, chunk: Bytes) -> Result<Option<Bytes>, FilterError> {
let mut body = Some(chunk);
self.run_step_body_filters(&mut body, false)?;
let emitted = self
.continuation
.extensions
.get_mut::<PendingStreamChunks>()
.map_or_else(VecDeque::new, PendingStreamChunks::drain_chunks);
self.pending_chunks.extend(emitted);
self.pending_chunks.extend(body.filter(|bytes| !bytes.is_empty()));
Ok(self.pending_chunks.pop_front())
}
async fn handle_filter_error(&mut self, error: FilterError) -> Result<Option<Bytes>, FilterError> {
warn!("filtered_subrequest: response body filter failed: {error}");
if let Some(upstream_body) = self.upstream.take() {
(*upstream_body).cancel().await;
}
self.continuation
.extensions
.insert(StreamTermination::new(StreamTerminationCause::Filter));
let completion = self.complete_step().map_err(|completion_error| -> FilterError {
format!(
"filtered_subrequest: response filter failed ({error}); completion also failed ({completion_error})"
)
.into()
})?;
self.finished = true;
self.deferred_completion_output = self.handled_completion_output(completion);
Ok(None)
}
fn handle_upstream_eof(&mut self) -> Result<Option<Bytes>, FilterError> {
let completion = self.complete_step()?;
self.finished = true;
self.deferred_completion_output = completion.filter(|bytes| !bytes.is_empty());
Ok(None)
}
async fn handle_upstream_error(
&mut self,
e: praxis_core::subrequest::SubRequestError,
) -> Result<Option<Bytes>, FilterError> {
if let Some(upstream_body) = self.upstream.take() {
(*upstream_body).cancel().await;
}
self.continuation
.extensions
.insert(StreamTermination::new(termination_cause(&e)));
let completion = self.complete_step()?;
self.finished = true;
self.deferred_completion_output = self.handled_completion_output(completion);
Ok(None)
}
fn handled_completion_output(&self, completion: Option<Bytes>) -> Option<Bytes> {
self.continuation
.extensions
.get::<StreamTermination>()
.is_some_and(StreamTermination::is_handled)
.then_some(completion)
.flatten()
.filter(|bytes| !bytes.is_empty())
}
}
#[async_trait]
impl StreamingResponseBody for FilteredStreamingBody {
#[expect(clippy::too_many_lines, reason = "pull loop applies deadlines and completion state")]
async fn next_chunk(&mut self) -> Result<Option<Bytes>, FilterError> {
if let Some(chunk) = self.pending_chunks.pop_front() {
return Ok(Some(chunk));
}
if self.finished {
return Ok(None);
}
loop {
let upstream = self
.upstream
.as_mut()
.ok_or_else(|| -> FilterError { "filtered_subrequest: upstream already consumed".to_owned().into() })?;
let remaining = self
.continuation
.step_deadline
.checked_duration_since(std::time::Instant::now())
.unwrap_or_default();
let next = if remaining.is_zero() {
Err(praxis_core::subrequest::SubRequestError::DeadlineExceeded)
} else {
tokio::time::timeout(remaining, upstream.next_chunk())
.await
.unwrap_or(Err(praxis_core::subrequest::SubRequestError::DeadlineExceeded))
};
match next {
Ok(Some(chunk)) => match self.handle_upstream_chunk(chunk) {
Ok(Some(bytes)) => return Ok(Some(bytes)),
Ok(None) => {},
Err(error) => return Box::pin(self.handle_filter_error(error)).await,
},
Ok(None) => return self.handle_upstream_eof(),
Err(e) => return Box::pin(self.handle_upstream_error(e)).await,
}
}
}
async fn suppress(&mut self) -> Result<(), FilterError> {
if !self.finished {
self.finished = true;
if let Some(upstream_body) = self.upstream.take() {
(*upstream_body).cancel().await;
}
self.complete_step()?;
}
Ok(())
}
async fn cancel(&mut self) {
if !self.finished {
self.finished = true;
if let Some(upstream_body) = self.upstream.take() {
(*upstream_body).cancel().await;
}
}
}
fn swap_extensions(&mut self, extensions: &mut RequestExtensions) {
self.exchange_extensions(extensions);
}
}
fn leftover_stream_read_timeout(ctx: &mut crate::filter::HttpFilterContext<'_>) -> Option<Duration> {
ctx.take_stream_read_timeout_cap()
.or_else(|| ctx.upstream.as_ref().and_then(|peer| peer.connection.read_timeout))
}
fn apply_leftover_read_timeout(body: &mut Option<Box<SubResponseBody>>, leftover: Option<Duration>) {
if let Some(timeout) = leftover
&& let Some(upstream) = body.as_mut()
{
upstream.cap_read_timeout(timeout);
}
}
fn leftover_stream_deadline(ctx: &mut crate::filter::HttpFilterContext<'_>) -> Option<std::time::Instant> {
ctx.take_stream_deadline_cap()
}
fn std_instant_to_tokio(deadline: std::time::Instant) -> tokio::time::Instant {
tokio::time::Instant::from_std(deadline)
}
fn apply_leftover_stream_deadline(body: &mut Option<Box<SubResponseBody>>, deadline: Option<std::time::Instant>) {
if let Some(deadline) = deadline
&& let Some(upstream) = body.as_mut()
{
upstream.cap_stream_deadline(std_instant_to_tokio(deadline));
}
}
pub(crate) struct CalloutStreamingBody {
inner: Option<FilteredStreamingBody>,
pending: VecDeque<Bytes>,
held_extensions: Option<RequestExtensions>,
deferred_error: Option<FilterError>,
emitted_bytes: usize,
max_response_bytes: usize,
finished: bool,
}
impl CalloutStreamingBody {
pub(crate) fn new(inner: FilteredStreamingBody, max_response_bytes: usize) -> Self {
Self {
inner: Some(inner),
pending: VecDeque::new(),
held_extensions: None,
deferred_error: None,
emitted_bytes: 0,
max_response_bytes,
finished: false,
}
}
fn checked(&mut self, chunk: Bytes) -> Result<Option<Bytes>, FilterError> {
let outcome = self.account(chunk);
if outcome.is_err() {
self.finished = true;
self.pending.clear();
}
outcome
}
fn account(&mut self, chunk: Bytes) -> Result<Option<Bytes>, FilterError> {
let total = self
.emitted_bytes
.checked_add(chunk.len())
.ok_or_else(|| -> FilterError { "filtered_subrequest: stream byte count overflow".into() })?;
if total > self.max_response_bytes {
return Err("filtered_subrequest: streaming response exceeds configured body limit"
.to_owned()
.into());
}
self.emitted_bytes = total;
Ok(Some(chunk))
}
fn drain_completion(&mut self) -> Result<(), FilterError> {
let (continuation, completion_output) = self
.inner
.take()
.ok_or_else(|| -> FilterError { "filtered_subrequest: streaming body already consumed".into() })?
.into_finished_parts();
let completion = continuation.into_completion();
self.held_extensions = Some(completion.extensions);
if let Some(termination) = completion.termination.as_ref().filter(|t| !t.is_handled()) {
let cause = termination.cause();
self.deferred_error =
Some(format!("filtered_subrequest: unhandled upstream stream termination: {cause:?}").into());
return Ok(());
}
self.pending.extend(completion.pending_chunks);
self.pending.extend(completion_output.filter(|bytes| !bytes.is_empty()));
self.finished = true;
Ok(())
}
}
#[async_trait]
impl StreamingResponseBody for CalloutStreamingBody {
async fn next_chunk(&mut self) -> Result<Option<Bytes>, FilterError> {
loop {
if let Some(chunk) = self.pending.pop_front() {
return self.checked(chunk);
}
if let Some(error) = self.deferred_error.take() {
self.finished = true;
self.inner = None;
return Err(error);
}
if self.finished {
return Ok(None);
}
let inner = self
.inner
.as_mut()
.ok_or_else(|| -> FilterError { "filtered_subrequest: streaming body has no active source".into() })?;
match inner.next_chunk().await? {
Some(chunk) => return self.checked(chunk),
None => self.drain_completion()?,
}
}
}
async fn suppress(&mut self) -> Result<(), FilterError> {
self.pending.clear();
self.deferred_error = None;
let outcome = if let Some(mut inner) = self.inner.take() {
let result = inner.suppress().await;
self.held_extensions = Some(inner.into_continuation().into_parent_extensions());
result
} else {
Ok(())
};
self.finished = true;
outcome
}
async fn cancel(&mut self) {
self.pending.clear();
self.deferred_error = None;
if let Some(mut inner) = self.inner.take() {
inner.cancel().await;
self.held_extensions = Some(inner.into_continuation().into_parent_extensions());
}
self.finished = true;
}
fn swap_extensions(&mut self, extensions: &mut RequestExtensions) {
if let Some(inner) = self.inner.as_mut() {
inner.swap_extensions(extensions);
} else if let Some(held) = self.held_extensions.as_mut() {
std::mem::swap(held, extensions);
}
}
}
#[cfg(test)]
#[expect(
clippy::expect_used,
clippy::too_many_lines,
clippy::unwrap_used,
reason = "streaming deadline integration setup"
)]
mod tests {
use std::{
collections::HashMap,
sync::Arc,
time::{Duration, Instant},
};
use async_trait::async_trait;
use bytes::Bytes;
use http::HeaderMap;
use praxis_core::subrequest::{StreamLimits, SubRequest, SubRequestClient, SubRequestError};
use super::{FilteredStreamingBody, std_instant_to_tokio};
struct ExpiredDeadlineFilter;
#[async_trait]
impl crate::HttpFilter for ExpiredDeadlineFilter {
fn name(&self) -> &'static str {
"expired_deadline_test_filter"
}
async fn on_request(
&self,
_ctx: &mut crate::HttpFilterContext<'_>,
) -> Result<crate::FilterAction, crate::FilterError> {
Ok(crate::FilterAction::Continue)
}
fn response_body_access(&self) -> crate::BodyAccess {
crate::BodyAccess::ReadOnly
}
fn on_response_body(
&self,
ctx: &mut crate::HttpFilterContext<'_>,
_body: &mut Option<Bytes>,
_end_of_stream: bool,
) -> Result<crate::FilterAction, crate::FilterError> {
ctx.cap_stream_deadline(Instant::now() - Duration::from_secs(1));
Ok(crate::FilterAction::Continue)
}
}
async fn spawn_stalling_backend() -> (std::net::SocketAddr, tokio::task::JoinHandle<()>) {
use tokio::io::AsyncWriteExt as _;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let handle = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0_u8; 4096];
drop(tokio::io::AsyncReadExt::read(&mut socket, &mut request).await);
socket
.write_all(b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n5\r\nhello\r\n")
.await
.unwrap();
socket.flush().await.unwrap();
tokio::time::sleep(Duration::from_secs(60)).await;
});
(addr, handle)
}
#[tokio::test]
async fn body_filter_deadline_is_applied_to_live_body() {
use pingora_core::upstreams::peer::HttpPeer;
let (addr, backend) = spawn_stalling_backend().await;
praxis_tls::provider::install();
let connector = praxis_core::subrequest::SubRequestConnector::new(1, None);
let client = SubRequestClient::new(connector);
let peer = HttpPeer::new(addr.to_string(), false, String::new());
let request = SubRequest {
method: http::Method::GET,
uri: http::Uri::from_static("/deadline-wire"),
headers: HeaderMap::new(),
body: Bytes::new(),
};
let response = Box::pin(client.send_streaming(
&peer,
&request,
Duration::from_secs(5),
StreamLimits {
idle_timeout: Duration::from_secs(30),
max_stream_duration: None,
max_total_bytes: None,
},
None,
))
.await
.unwrap();
let mut registry = crate::FilterRegistry::with_builtins();
registry
.register(
"expired_deadline_test_filter",
crate::FilterFactory::Http(Arc::new(|_| Ok(Box::new(ExpiredDeadlineFilter)))),
)
.unwrap();
let mut entries = vec![praxis_core::config::FilterEntry {
branch_chains: None,
conditions: vec![],
filter_type: "expired_deadline_test_filter".into(),
config: serde_yaml::Value::Null,
name: None,
response_conditions: vec![],
failure_mode: praxis_core::config::FailureMode::default(),
}];
let pipeline = Arc::new(crate::FilterPipeline::build(&mut entries, ®istry).unwrap());
let continuation = super::super::continuation::FilteredSubrequestContinuation {
pipeline,
request_snapshot: crate::Request {
headers: HeaderMap::new(),
method: http::Method::GET,
uri: http::Uri::from_static("/deadline-wire"),
},
response_snapshot: crate::Response {
headers: HeaderMap::new(),
status: http::StatusCode::OK,
},
extensions: crate::RequestExtensions::default(),
filter_state: HashMap::new(),
filter_results: HashMap::new(),
filter_metadata: HashMap::new(),
structured_metadata: HashMap::new(),
executed_filter_indices: vec![true],
body_done_indices: vec![false],
response_body_bytes: 0,
response_body_mode: crate::BodyMode::Stream,
completed: false,
client_addr: None,
downstream_tls: false,
request_start: Instant::now(),
step_deadline: Instant::now() + Duration::from_secs(30),
peer_identity: None,
};
let mut filtered = FilteredStreamingBody::new(Box::new(response.body), continuation);
let mut body = Some(Bytes::from_static(b"first chunk"));
filtered.run_step_body_filters(&mut body, false).unwrap();
let err = filtered
.upstream
.as_mut()
.expect("the live upstream body must remain attached")
.next_chunk()
.await
.unwrap_err();
drop(filtered);
backend.abort();
assert!(
matches!(err, SubRequestError::DeadlineExceeded),
"the body filter's expired absolute deadline must reach the live body: {err}"
);
}
#[test]
fn std_instant_to_tokio_preserves_a_future_deadline() {
let deadline = Instant::now() + Duration::from_secs(1);
let converted = std_instant_to_tokio(deadline);
assert!(
converted > tokio::time::Instant::now(),
"a future standard deadline must remain in the future after conversion"
);
assert_eq!(
converted.into_std(),
deadline,
"direct conversion must preserve the absolute deadline"
);
}
#[test]
fn std_instant_to_tokio_clamps_an_expired_deadline() {
let deadline = Instant::now() - Duration::from_secs(1);
let converted = std_instant_to_tokio(deadline);
assert!(
converted <= tokio::time::Instant::now(),
"an expired standard deadline must not become a future Tokio deadline"
);
assert_eq!(
converted.into_std(),
deadline,
"direct conversion must preserve an expired absolute deadline"
);
}
}