use axum::body::{Body, Bytes};
use axum::response::{IntoResponse, Response};
use futures_core::Stream;
use futures_util::StreamExt as _;
use http::HeaderValue;
use http::header::CONTENT_TYPE;
use serde::Serialize;
use toolkit_canonical_errors::Problem;
use toolkit_contract::runtime::multipart::MAX_ACCUMULATED_BYTES;
const MAX_BOUNDARY_LEN: usize = 70;
const PROBLEM_CONTENT_TYPE: &str = "application/problem+json";
const MAX_PART_BYTES: usize = MAX_ACCUMULATED_BYTES;
pub struct MultipartJsonStream<S> {
stream: S,
boundary: String,
content_type: HeaderValue,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct BoundaryError {
boundary: String,
reason: &'static str,
}
impl BoundaryError {
#[must_use]
pub fn reason(&self) -> &str {
self.reason
}
}
impl std::fmt::Display for BoundaryError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(
f,
"invalid multipart/mixed boundary \"{}\": {}",
self.boundary.escape_debug(),
self.reason
)
}
}
impl std::error::Error for BoundaryError {}
impl<S> MultipartJsonStream<S> {
#[must_use]
pub fn new(stream: S) -> Self {
Self::with_validated_boundary(stream, generated_boundary())
}
pub fn with_boundary(stream: S, boundary: impl Into<String>) -> Result<Self, BoundaryError> {
let boundary = boundary.into();
match validate_boundary(&boundary) {
Ok(()) => Ok(Self::with_validated_boundary(stream, boundary)),
Err(reason) => Err(BoundaryError { boundary, reason }),
}
}
fn with_validated_boundary(stream: S, boundary: String) -> Self {
if let Some(content_type) = content_type_for(&boundary) {
return Self {
stream,
boundary,
content_type,
};
}
tracing::error!(
boundary = ?boundary,
"validated multipart/mixed boundary did not form a legal header value; substituting a known-good boundary"
);
Self {
stream,
boundary: FALLBACK_BOUNDARY.to_owned(),
content_type: HeaderValue::from_static(FALLBACK_CONTENT_TYPE),
}
}
#[must_use]
pub fn boundary(&self) -> &str {
&self.boundary
}
}
impl<S, T, E> IntoResponse for MultipartJsonStream<S>
where
S: Stream<Item = Result<T, E>> + Send + 'static,
T: Serialize,
E: Into<Problem>,
{
fn into_response(self) -> Response {
let mut response = Response::new(Body::from_stream(frame(self.stream, self.boundary)));
response
.headers_mut()
.insert(CONTENT_TYPE, self.content_type);
response
}
}
struct FramerState<S> {
stream: S,
boundary: String,
finished: bool,
}
fn frame<S, T, E>(
stream: S,
boundary: String,
) -> impl Stream<Item = Result<Bytes, std::io::Error>> + Send + 'static
where
S: Stream<Item = Result<T, E>> + Send + 'static,
T: Serialize,
E: Into<Problem>,
{
let state = FramerState {
stream: Box::pin(stream),
boundary,
finished: false,
};
futures_util::stream::unfold(state, |mut state| async move {
if state.finished {
return None;
}
let Some(item) = state.stream.next().await else {
state.finished = true;
let close = format!("--{}--\r\n", state.boundary);
return Some((Ok(Bytes::from(close)), state));
};
let value = match item {
Ok(value) => value,
Err(e) => {
state.finished = true;
let mut problem: Problem = e.into();
tracing::warn!(
status = ?problem.status,
error_code = ?problem.error_code,
"multipart/mixed stream ended with a domain error; framing it as a typed error part"
);
let Some(bytes) = encode_error_part_within_limit(&state.boundary, &mut problem)
else {
tracing::error!(
status = ?problem.status,
"multipart/mixed error part exceeds the maximum part size even after trimming, or would not serialize; aborting the response body"
);
return Some((
Err(std::io::Error::other(
"multipart/mixed error part exceeds the maximum part size even after trimming",
)),
state,
));
};
return Some((Ok(bytes), state));
}
};
match serde_json::to_vec(&value) {
Ok(json) if json.len() > MAX_PART_BYTES => {
tracing::error!(
item_type = std::any::type_name::<T>(),
item_bytes = json.len(),
max_bytes = MAX_PART_BYTES,
"multipart/mixed stream item exceeds the maximum part size; aborting the response body"
);
state.finished = true;
Some((
Err(std::io::Error::other(format!(
"multipart/mixed stream item is {} bytes, exceeding the maximum part size of {MAX_PART_BYTES} bytes",
json.len()
))),
state,
))
}
Ok(json) => {
let bytes = encode_part(&state.boundary, &json);
Some((Ok(bytes), state))
}
Err(e) => {
tracing::error!(
item_type = std::any::type_name::<T>(),
error = %e,
"failed to serialize a multipart/mixed stream item; aborting the response body"
);
state.finished = true;
Some((
Err(std::io::Error::other(format!(
"failed to serialize a multipart/mixed stream item: {e}"
))),
state,
))
}
}
})
}
fn encode_part(boundary: &str, json: &[u8]) -> Bytes {
let header = format!(
"--{boundary}\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n",
json.len()
);
let mut out = Vec::with_capacity(header.len() + json.len() + 2);
out.extend_from_slice(header.as_bytes());
out.extend_from_slice(json);
out.extend_from_slice(b"\r\n");
Bytes::from(out)
}
fn encode_error_part_and_close(boundary: &str, problem_json: &[u8]) -> Bytes {
let header = format!(
"--{boundary}\r\nContent-Type: {PROBLEM_CONTENT_TYPE}\r\nContent-Length: {}\r\n\r\n",
problem_json.len()
);
let close = format!("\r\n--{boundary}--\r\n");
let mut out = Vec::with_capacity(header.len() + problem_json.len() + close.len());
out.extend_from_slice(header.as_bytes());
out.extend_from_slice(problem_json);
out.extend_from_slice(close.as_bytes());
Bytes::from(out)
}
fn encode_error_part_within_limit(boundary: &str, problem: &mut Problem) -> Option<Bytes> {
let full = serde_json::to_vec(problem)
.ok()
.filter(|json| json.len() <= MAX_PART_BYTES);
if let Some(json) = full {
return Some(encode_error_part_and_close(boundary, &json));
}
problem.detail = String::new();
problem.context = serde_json::Value::Null;
let trimmed = serde_json::to_vec(problem)
.ok()
.filter(|json| json.len() <= MAX_PART_BYTES)?;
Some(encode_error_part_and_close(boundary, &trimmed))
}
fn generated_boundary() -> String {
uuid::Uuid::now_v7().simple().to_string()
}
fn content_type_for(boundary: &str) -> Option<HeaderValue> {
HeaderValue::from_str(&format!("multipart/mixed; boundary={boundary}")).ok()
}
const FALLBACK_BOUNDARY: &str = "0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f";
const FALLBACK_CONTENT_TYPE: &str = "multipart/mixed; boundary=0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f0f";
fn validate_boundary(boundary: &str) -> Result<(), &'static str> {
if boundary.is_empty() {
return Err("boundary must not be empty");
}
if boundary.len() > MAX_BOUNDARY_LEN {
return Err("boundary exceeds 70 characters");
}
if !boundary
.bytes()
.all(|b| b.is_ascii_alphanumeric() || b"'+-._".contains(&b))
{
return Err("boundary contains a character outside the HTTP token set");
}
Ok(())
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
#[allow(clippy::unwrap_used)]
mod tests {
use super::*;
use axum::body::to_bytes;
use serde::Serialize;
use toolkit_canonical_errors::CanonicalError;
#[derive(Serialize)]
struct Item {
id: u32,
}
struct Unserializable;
impl Serialize for Unserializable {
fn serialize<S: serde::Serializer>(&self, _: S) -> Result<S::Ok, S::Error> {
Err(serde::ser::Error::custom("nope"))
}
}
#[tokio::test]
async fn frames_one_part_per_item_with_a_content_length() {
let items = futures_util::stream::iter(vec![
Ok::<_, CanonicalError>(Item { id: 1 }),
Ok(Item { id: 2 }),
]);
let response = MultipartJsonStream::with_boundary(items, "BOUND")
.unwrap()
.into_response();
assert_eq!(
response.headers().get(CONTENT_TYPE).unwrap(),
"multipart/mixed; boundary=BOUND"
);
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
assert_eq!(
String::from_utf8(body.to_vec()).unwrap(),
"--BOUND\r\nContent-Type: application/json\r\nContent-Length: 8\r\n\r\n{\"id\":1}\r\n\
--BOUND\r\nContent-Type: application/json\r\nContent-Length: 8\r\n\r\n{\"id\":2}\r\n\
--BOUND--\r\n"
);
}
#[tokio::test]
async fn an_empty_stream_is_just_the_close_delimiter() {
let items = futures_util::stream::iter(Vec::<Result<Item, CanonicalError>>::new());
let response = MultipartJsonStream::with_boundary(items, "BOUND")
.unwrap()
.into_response();
let body = to_bytes(response.into_body(), usize::MAX).await.unwrap();
assert_eq!(String::from_utf8(body.to_vec()).unwrap(), "--BOUND--\r\n");
}
#[tokio::test]
async fn an_unserializable_item_aborts_the_body_rather_than_ending_it() {
let items = futures_util::stream::iter(vec![Ok::<_, CanonicalError>(Unserializable)]);
let response = MultipartJsonStream::with_boundary(items, "BOUND")
.unwrap()
.into_response();
let result = to_bytes(response.into_body(), usize::MAX).await;
assert!(
result.is_err(),
"the body must abort, not end gracefully; got {:?}",
result.map(|b| String::from_utf8_lossy(&b).into_owned())
);
}
#[tokio::test]
async fn an_err_item_becomes_a_typed_error_part_then_a_clean_close() {
let items = futures_util::stream::iter(vec![
Ok(Item { id: 1 }),
Err(CanonicalError::internal("mid-stream boom").create()),
Ok(Item { id: 2 }),
]);
let response = MultipartJsonStream::with_boundary(items, "BOUND")
.unwrap()
.into_response();
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("an error item ends the body cleanly, it must not abort");
let text = String::from_utf8(body.to_vec()).unwrap();
assert!(
text.starts_with(
"--BOUND\r\nContent-Type: application/json\r\nContent-Length: 8\r\n\r\n{\"id\":1}\r\n"
),
"expected the data part first; got:\n{text}"
);
assert!(
text.contains("Content-Type: application/problem+json"),
"expected a problem+json error part; got:\n{text}"
);
assert!(
text.ends_with("--BOUND--\r\n"),
"expected a clean close; got:\n{text}"
);
assert!(
!text.contains("{\"id\":2}"),
"an error is terminal; later items must not be framed; got:\n{text}"
);
}
#[tokio::test]
async fn an_oversized_error_problem_is_trimmed_not_aborted() {
let huge = "a".repeat(MAX_PART_BYTES + 4096);
let items = futures_util::stream::iter(vec![Err::<Item, CanonicalError>(
CanonicalError::internal(huge).create(),
)]);
let response = MultipartJsonStream::with_boundary(items, "BOUND")
.unwrap()
.into_response();
let body = to_bytes(response.into_body(), usize::MAX)
.await
.expect("an oversized error Problem must be trimmed and framed, not aborted");
let text = String::from_utf8_lossy(&body);
assert!(
text.contains("Content-Type: application/problem+json"),
"expected a problem+json error part; got {} bytes",
body.len()
);
assert!(
text.ends_with("--BOUND--\r\n"),
"expected a clean close; got {} bytes",
body.len()
);
assert!(
body.len() < MAX_PART_BYTES,
"the oversized detail must be trimmed away; body was {} bytes",
body.len()
);
}
#[tokio::test]
async fn an_oversized_item_aborts_the_body_rather_than_emitting_an_unreadable_part() {
let oversized = "a".repeat(MAX_PART_BYTES);
let items = futures_util::stream::iter(vec![Ok::<_, CanonicalError>(oversized)]);
let response = MultipartJsonStream::with_boundary(items, "BOUND")
.unwrap()
.into_response();
let result = to_bytes(response.into_body(), usize::MAX).await;
assert!(
result.is_err(),
"an oversized item must abort the body, not emit an unreadable part; got {:?}",
result.map(|b| b.len())
);
}
#[test]
fn the_emitted_content_type_carries_the_framing_boundary() {
let framer = MultipartJsonStream::with_boundary(
futures_util::stream::iter(Vec::<Result<Item, CanonicalError>>::new()),
"abc123",
)
.unwrap();
let expected = format!("multipart/mixed; boundary={}", framer.boundary());
let response = framer.into_response();
assert_eq!(response.headers().get(CONTENT_TYPE).unwrap(), &expected);
}
#[test]
fn the_fallback_boundary_and_header_stay_consistent() {
assert_eq!(
FALLBACK_CONTENT_TYPE,
format!("multipart/mixed; boundary={FALLBACK_BOUNDARY}")
);
assert!(validate_boundary(FALLBACK_BOUNDARY).is_ok());
assert!(content_type_for(FALLBACK_BOUNDARY).is_some());
}
#[test]
fn a_generated_boundary_is_used_when_none_is_given() {
let items = futures_util::stream::iter(Vec::<Result<Item, CanonicalError>>::new());
let framer = MultipartJsonStream::new(items);
assert!(validate_boundary(framer.boundary()).is_ok());
assert!(!framer.boundary().is_empty());
}
#[test]
fn an_illegal_boundary_is_rejected_not_silently_substituted() {
let items = futures_util::stream::iter(Vec::<Result<Item, CanonicalError>>::new());
let Err(err) = MultipartJsonStream::with_boundary(items, "not\r\nlegal") else {
panic!("an illegal boundary must be rejected");
};
assert!(
err.to_string().contains("not\\r\\nlegal"),
"message should escape control chars: {err}"
);
let framer = MultipartJsonStream::new(futures_util::stream::iter(Vec::<
Result<Item, CanonicalError>,
>::new()));
assert!(validate_boundary(framer.boundary()).is_ok());
}
#[test]
fn boundary_validation_requires_an_rfc_2046_http_token() {
assert!(validate_boundary("abcABC012'+-._").is_ok());
assert!(validate_boundary("").is_err());
assert!(validate_boundary(&"a".repeat(MAX_BOUNDARY_LEN + 1)).is_err());
assert!(validate_boundary("has space").is_err());
assert!(validate_boundary("a/b").is_err());
assert!(validate_boundary("a:b=c?").is_err());
assert!(validate_boundary("(paren)").is_err());
assert!(validate_boundary("comma,d").is_err());
assert!(validate_boundary("semi;colon").is_err());
assert!(validate_boundary("quote\"d").is_err());
}
}