use std::{collections::HashMap, sync::Arc};
use bytes::Bytes;
use praxis_core::{config::ABSOLUTE_MAX_BODY_BYTES, connectivity::Upstream};
use praxis_filter::{BodyBuffer, BodyMode, RequestExtensions};
use crate::http::pingora::context::PingoraRequestCtx;
pub(super) fn clamp_body_mode_to_ceiling(mode: BodyMode, baseline: BodyMode) -> BodyMode {
let ceiling = match baseline {
BodyMode::StreamBuffer { max_bytes: Some(v) } | BodyMode::SizeLimit { max_bytes: v } => Some(v),
_ => None,
};
match (mode, ceiling) {
(BodyMode::StreamBuffer { max_bytes }, Some(limit)) => BodyMode::StreamBuffer {
max_bytes: Some(max_bytes.map_or(limit, |v| v.min(limit))),
},
(BodyMode::SizeLimit { max_bytes }, Some(limit)) => BodyMode::SizeLimit {
max_bytes: max_bytes.min(limit),
},
(m, None | Some(_)) => m,
}
}
pub(super) fn check_body_size_limit(body: Option<&Bytes>, accumulated_bytes: &mut u64, max_bytes: usize) -> bool {
if let Some(chunk) = body {
let chunk_len = chunk.len() as u64;
*accumulated_bytes += chunk_len;
let limit = max_bytes as u64;
return *accumulated_bytes > limit;
}
false
}
pub(super) fn exceeds_body_ceiling(ceiling: Option<usize>, counted: u64, chunk: Option<&Bytes>) -> bool {
let chunk_len = chunk.map_or(0, Bytes::len) as u64;
ceiling.is_some_and(|max| counted.saturating_add(chunk_len) > max as u64)
}
pub(super) fn accumulate_stream_buffer(
body: &mut Option<Bytes>,
body_buffer: &mut Option<BodyBuffer>,
end_of_stream: bool,
max_bytes: Option<usize>,
) -> bool {
if let Some(chunk) = &*body {
let limit = max_bytes.unwrap_or(ABSOLUTE_MAX_BODY_BYTES);
let buf = body_buffer.get_or_insert_with(|| BodyBuffer::new(limit));
if buf.push(chunk.clone()).is_err() {
return true;
}
}
if end_of_stream {
tracing::trace!("stream buffer: freezing accumulated body before pipeline at EOS");
*body = body_buffer.take().map(BodyBuffer::freeze);
} else {
tracing::trace!("stream buffer: filters see the original chunk");
}
false
}
#[expect(
clippy::fn_params_excessive_bools,
reason = "mirrors the caller's existing condition flags"
)]
pub(super) fn suppress_stream_buffer_chunk(
body: &mut Option<Bytes>,
is_stream_buffer: bool,
released: bool,
end_of_stream: bool,
) {
if is_stream_buffer && !released && !end_of_stream {
*body = None;
}
}
pub(super) fn release_stream_buffer(
body: &mut Option<Bytes>,
is_stream_buffer: bool,
released: &mut bool,
body_buffer: &mut Option<BodyBuffer>,
end_of_stream: bool,
) {
if is_stream_buffer && !*released {
*released = true;
if !end_of_stream {
*body = body_buffer.take().map(BodyBuffer::freeze);
}
}
}
pub(super) struct BodyFilterOutput {
pub(super) cluster: Option<Arc<str>>,
pub(super) upstream: Option<Upstream>,
pub(super) extensions: RequestExtensions,
pub(super) filter_metadata: HashMap<String, String>,
pub(super) filter_state: HashMap<usize, Box<dyn std::any::Any + Send + Sync>>,
pub(super) executed_branch_filters: Vec<bool>,
pub(super) executed_filter_indices: Vec<bool>,
pub(super) body_done_indices: Vec<bool>,
pub(super) attempted_endpoints: Vec<Arc<str>>,
}
impl BodyFilterOutput {
pub(super) fn take_from(fctx: &mut praxis_filter::HttpFilterContext<'_>) -> Self {
Self {
cluster: fctx.cluster.take(),
upstream: fctx.upstream.take(),
extensions: std::mem::take(&mut fctx.extensions),
filter_metadata: std::mem::take(&mut fctx.filter_metadata),
filter_state: std::mem::take(&mut fctx.filter_state),
executed_branch_filters: std::mem::take(&mut fctx.executed_branch_filters),
executed_filter_indices: std::mem::take(&mut fctx.executed_filter_indices),
body_done_indices: std::mem::take(&mut fctx.body_done_indices),
attempted_endpoints: std::mem::take(&mut fctx.attempted_endpoints),
}
}
pub(super) fn write_back(self, ctx: &mut PingoraRequestCtx) {
ctx.cluster = self.cluster;
ctx.upstream = self.upstream;
ctx.extensions = self.extensions;
ctx.filter_metadata = self.filter_metadata;
ctx.filter_state = self.filter_state;
ctx.cached_executed_branch_filters = self.executed_branch_filters;
ctx.cached_executed_filter_indices = self.executed_filter_indices;
ctx.cached_body_done_indices = self.body_done_indices;
ctx.attempted_endpoints = self.attempted_endpoints;
}
}