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 result = {
let cont = &mut self.continuation;
let mut ctx = crate::filter::HttpFilterContext {
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_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);
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);
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);
}
}
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);
}
}
}