use crate::domain::entities::Frame;
use async_stream::try_stream;
use axum::{
http::{HeaderMap, StatusCode, header},
response::Response,
};
use futures::{Stream, StreamExt};
use headers_accept::Accept;
use mediatype::{MediaType, MediaTypeBuf, Name, names};
use std::str::FromStr;
#[derive(Debug, Clone, Copy)]
pub enum StreamFormat {
Json,
NdJson,
ServerSentEvents,
Binary,
}
const MAX_ACCEPT_ENTRIES: usize = 16;
const X_NDJSON: Name<'static> = Name::new_unchecked("x-ndjson");
static SUPPORTED_MEDIA_TYPES: [(MediaType<'static>, StreamFormat); 4] = [
(
MediaType::new(names::APPLICATION, names::JSON),
StreamFormat::Json,
),
(
MediaType::new(names::TEXT, names::EVENT_STREAM),
StreamFormat::ServerSentEvents,
),
(
MediaType::new(names::APPLICATION, X_NDJSON),
StreamFormat::NdJson,
),
(
MediaType::new(names::APPLICATION, names::OCTET_STREAM),
StreamFormat::Binary,
),
];
fn is_supported_media_range(media_range: &str) -> bool {
!media_range.contains('*')
|| media_range.eq_ignore_ascii_case("*/*")
|| media_range.eq_ignore_ascii_case("application/*")
}
fn format_q(q: f32) -> String {
if q <= 0.0 {
return "0.000".to_string();
}
let milli = ((q * 1000.0).round() as u32).max(1);
if milli >= 1000 {
"1.000".to_string()
} else {
format!("0.{milli:03}")
}
}
impl StreamFormat {
pub fn from_accept_header(headers: &HeaderMap) -> Self {
let Some(accept) = headers.get(header::ACCEPT) else {
return Self::Json;
};
let Ok(accept_str) = accept.to_str() else {
return Self::Json;
};
let mut sanitized_entries: Vec<String> = Vec::new();
for entry in accept_str.split(',').take(MAX_ACCEPT_ENTRIES) {
let mut parts = entry.split(';');
let media_range = parts.next().unwrap_or("").trim();
if media_range.is_empty() || !is_supported_media_range(media_range) {
continue;
}
let mut q_str: Option<&str> = None;
for param in parts {
let mut kv = param.splitn(2, '=');
let name = kv.next().unwrap_or("").trim();
if name.eq_ignore_ascii_case("q") {
q_str = Some(kv.next().unwrap_or("").trim());
break;
}
}
let sanitized = match q_str {
None => media_range.to_string(),
Some(raw_q) => {
let Ok(q) = raw_q.parse::<f32>() else {
continue;
};
if !q.is_finite() {
continue;
}
format!("{media_range};q={}", format_q(q.clamp(0.0, 1.0)))
}
};
if MediaTypeBuf::from_str(&sanitized).is_err() {
continue;
}
sanitized_entries.push(sanitized);
}
if sanitized_entries.is_empty() {
return Self::Json;
}
let Ok(accept) = Accept::from_str(&sanitized_entries.join(",")) else {
return Self::Json;
};
let Some(best) = accept.negotiate(SUPPORTED_MEDIA_TYPES.iter().map(|(mt, _)| mt)) else {
return Self::Json;
};
SUPPORTED_MEDIA_TYPES
.iter()
.find(|(mt, _)| mt == best)
.map_or(Self::Json, |(_, format)| *format)
}
pub fn content_type(&self) -> &'static str {
match self {
Self::Json => "application/json",
Self::NdJson => "application/x-ndjson",
Self::ServerSentEvents => "text/event-stream",
Self::Binary => "application/octet-stream",
}
}
}
fn format_batch_owned(
frames: &[Frame],
format: StreamFormat,
) -> Result<Vec<u8>, StreamTransportError> {
match format {
StreamFormat::Json | StreamFormat::NdJson => {
let mut out = Vec::new();
for frame in frames {
out.extend_from_slice(&sonic_rs::to_vec(frame)?);
out.push(b'\n');
}
Ok(out)
}
StreamFormat::ServerSentEvents => {
let mut out = Vec::new();
for frame in frames {
out.extend_from_slice(b"data: ");
out.extend_from_slice(&sonic_rs::to_vec(frame)?);
out.extend_from_slice(b"\n\n");
}
Ok(out)
}
StreamFormat::Binary => Ok(sonic_rs::to_vec(frames)?),
}
}
pub struct BatchFrameStream<S> {
inner: S,
format: StreamFormat,
batch_size: usize,
}
impl<S> BatchFrameStream<S>
where
S: Stream<Item = Frame> + Unpin + Send + 'static,
{
pub fn new(stream: S, format: StreamFormat, batch_size: usize) -> Self {
Self {
inner: stream,
format,
batch_size,
}
}
pub fn content_type(&self) -> &'static str {
match self.format {
StreamFormat::Json => "application/x-ndjson",
other => other.content_type(),
}
}
pub fn into_stream(
self,
) -> impl Stream<Item = Result<Vec<u8>, StreamTransportError>> + Send + 'static {
let Self {
inner,
format,
batch_size,
} = self;
try_stream! {
let mut batch: Vec<Frame> = Vec::with_capacity(batch_size);
futures::pin_mut!(inner);
while let Some(frame) = inner.next().await {
batch.push(frame);
if batch.len() >= batch_size {
let bytes = format_batch_owned(&batch, format)?;
batch.clear();
yield bytes;
}
}
if !batch.is_empty() {
let bytes = format_batch_owned(&batch, format)?;
yield bytes;
}
}
}
}
#[derive(Debug, thiserror::Error)]
pub enum StreamTransportError {
#[error("Serialization error: {0}")]
Serialization(#[from] sonic_rs::Error),
#[error("IO error: {0}")]
Io(String),
#[error("Buffer overflow")]
BufferOverflow,
#[error("Stream closed")]
StreamClosed,
}
pub fn create_streaming_response<S>(
stream: S,
format: StreamFormat,
) -> Result<Response, StreamTransportError>
where
S: Stream<Item = Result<Vec<u8>, StreamTransportError>> + Send + 'static,
{
let body = axum::body::Body::from_stream(stream);
let mut response = Response::builder()
.status(StatusCode::OK)
.header(header::CONTENT_TYPE, format.content_type())
.header(header::CACHE_CONTROL, "no-cache");
if let StreamFormat::ServerSentEvents = format {
response = response.header("X-Accel-Buffering", "no");
}
response
.body(body)
.map_err(|e| StreamTransportError::Io(e.to_string()))
}
pub fn create_streaming_response_with_content_type<S>(
stream: S,
content_type: &str,
) -> Result<Response, StreamTransportError>
where
S: Stream<Item = Result<Vec<u8>, StreamTransportError>> + Send + 'static,
{
let body = axum::body::Body::from_stream(stream);
Response::builder()
.status(StatusCode::OK)
.header(header::CONTENT_TYPE, content_type)
.header(header::CACHE_CONTROL, "no-cache")
.body(body)
.map_err(|e| StreamTransportError::Io(e.to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::domain::entities::Frame;
use crate::domain::value_objects::{JsonData, StreamId};
use axum::http::header;
use futures::StreamExt;
use futures::stream;
use std::pin::Pin;
use std::task::{Context, Poll};
fn make_skeleton_frame() -> Frame {
Frame::skeleton(StreamId::new(), 1, JsonData::Null)
}
struct PendingThenReady<I: Iterator> {
iter: I,
pending_remaining: usize,
pending_per_item: usize,
done: bool,
}
impl<I: Iterator> PendingThenReady<I> {
fn new(iter: I, pending_per_item: usize) -> Self {
Self {
iter,
pending_remaining: pending_per_item,
pending_per_item,
done: false,
}
}
}
impl<I: Iterator + Unpin> Stream for PendingThenReady<I> {
type Item = I::Item;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
if self.done {
return Poll::Ready(None);
}
if self.pending_remaining > 0 {
self.pending_remaining -= 1;
cx.waker().wake_by_ref();
return Poll::Pending;
}
match self.iter.next() {
Some(item) => {
self.pending_remaining = self.pending_per_item;
Poll::Ready(Some(item))
}
None => {
self.done = true;
Poll::Ready(None)
}
}
}
}
#[test]
fn test_stream_format_detection() {
let mut headers = HeaderMap::new();
headers.insert(header::ACCEPT, "text/event-stream".parse().unwrap());
let format = StreamFormat::from_accept_header(&headers);
assert!(matches!(format, StreamFormat::ServerSentEvents));
}
#[tokio::test]
async fn test_batch_frame_stream_multiple_batches() {
let frames: Vec<Frame> = (0..5).map(|_| make_skeleton_frame()).collect();
let frame_stream = stream::iter(frames);
let batch_stream = BatchFrameStream::new(frame_stream, StreamFormat::Json, 2);
let collected: Vec<Result<Vec<u8>, StreamTransportError>> =
batch_stream.into_stream().collect().await;
assert_eq!(
collected.len(),
3,
"expected 3 batches for 5 frames with batch_size=2"
);
let mut total_objects = 0usize;
for result in &collected {
let batch_bytes = result.as_ref().expect("batch should not error");
let batch_str = std::str::from_utf8(batch_bytes).expect("uncompressed batch is UTF-8");
for line in batch_str.lines() {
if line.is_empty() {
continue;
}
let parsed: serde_json::Value =
serde_json::from_str(line).expect("each line must be valid JSON");
assert!(
parsed.is_object(),
"each line must be a JSON object (NDJSON-of-objects), got: {line}"
);
total_objects += 1;
}
}
assert_eq!(
total_objects, 5,
"total parsed objects across all batches must equal 5"
);
}
#[test]
fn test_batch_stream_emits_only_full_batches_under_pending() {
tokio_test::block_on(async {
let frames: Vec<Frame> = (0..6).map(|_| make_skeleton_frame()).collect();
let inner = PendingThenReady::new(frames.into_iter(), 2);
let batch = BatchFrameStream::new(inner, StreamFormat::Json, 3);
let collected: Vec<_> = batch.into_stream().collect().await;
assert_eq!(
collected.len(),
2,
"6 frames at batch_size=3 must yield exactly 2 batches"
);
for r in collected {
assert!(r.is_ok());
}
});
}
#[tokio::test]
async fn test_batch_stream_ndjson_objects_per_line() {
let make_frames = || -> Vec<Frame> { (0..3).map(|_| make_skeleton_frame()).collect() };
let result_json: Vec<_> =
BatchFrameStream::new(stream::iter(make_frames()), StreamFormat::Json, 10)
.into_stream()
.collect()
.await;
assert_eq!(result_json.len(), 1);
let json_bytes = result_json[0].as_ref().unwrap();
let json_str = std::str::from_utf8(json_bytes).unwrap();
for line in json_str.lines() {
if line.is_empty() {
continue;
}
let v: serde_json::Value = serde_json::from_str(line).unwrap();
assert!(v.is_object(), "Json format: each line must be an object");
}
let result_ndjson: Vec<_> =
BatchFrameStream::new(stream::iter(make_frames()), StreamFormat::NdJson, 10)
.into_stream()
.collect()
.await;
assert_eq!(result_ndjson.len(), 1);
let ndjson_bytes = result_ndjson[0].as_ref().unwrap();
let ndjson_str = std::str::from_utf8(ndjson_bytes).unwrap();
for line in ndjson_str.lines() {
if line.is_empty() {
continue;
}
let v: serde_json::Value = serde_json::from_str(line).unwrap();
assert!(v.is_object(), "NdJson format: each line must be an object");
}
let json_count = json_str.lines().filter(|l| !l.is_empty()).count();
let ndjson_count = ndjson_str.lines().filter(|l| !l.is_empty()).count();
assert_eq!(
json_count, ndjson_count,
"Json and NdJson must produce the same object count"
);
let result_sse: Vec<_> = BatchFrameStream::new(
stream::iter(make_frames()),
StreamFormat::ServerSentEvents,
10,
)
.into_stream()
.collect()
.await;
assert_eq!(result_sse.len(), 1);
let sse_bytes = result_sse[0].as_ref().unwrap();
let sse_str = std::str::from_utf8(sse_bytes).unwrap();
let sse_frames: Vec<&str> = sse_str.split("\n\n").filter(|s| !s.is_empty()).collect();
assert_eq!(sse_frames.len(), 3);
for frame_str in sse_frames {
assert!(frame_str.starts_with("data: "));
let json_part = &frame_str["data: ".len()..];
let v: serde_json::Value = serde_json::from_str(json_part).unwrap();
assert!(v.is_object());
}
let result_binary: Vec<_> =
BatchFrameStream::new(stream::iter(make_frames()), StreamFormat::Binary, 10)
.into_stream()
.collect()
.await;
assert_eq!(result_binary.len(), 1);
let binary_bytes = result_binary[0].as_ref().unwrap();
let v: serde_json::Value = serde_json::from_slice(binary_bytes).unwrap();
assert!(v.is_array());
assert_eq!(v.as_array().unwrap().len(), 3);
}
#[tokio::test]
async fn test_create_streaming_response_with_content_type_uses_explicit_type() {
let frames: Vec<Frame> = (0..2).map(|_| make_skeleton_frame()).collect();
let batch = BatchFrameStream::new(stream::iter(frames), StreamFormat::Json, 10);
let expected_ct = batch.content_type();
assert_eq!(
expected_ct, "application/x-ndjson",
"BatchFrameStream with Json format must report application/x-ndjson"
);
let response =
create_streaming_response_with_content_type(batch.into_stream(), expected_ct)
.expect("response must be built");
let ct = response
.headers()
.get(header::CONTENT_TYPE)
.expect("Content-Type header must be present")
.to_str()
.unwrap();
assert_eq!(ct, "application/x-ndjson");
}
#[tokio::test]
async fn test_create_streaming_response_uses_format_content_type() {
let frames: Vec<Frame> = (0..1).map(|_| make_skeleton_frame()).collect();
let batch = BatchFrameStream::new(stream::iter(frames), StreamFormat::Json, 10);
let response = create_streaming_response(batch.into_stream(), StreamFormat::Json)
.expect("response must be built");
let ct = response
.headers()
.get(header::CONTENT_TYPE)
.expect("Content-Type header must be present")
.to_str()
.unwrap();
assert_eq!(ct, "application/json");
}
#[test]
fn test_sonic_rs_matches_serde_json_semantics() {
let cases: Vec<(&str, serde_json::Value)> = vec![
("empty_object", serde_json::json!({})),
("empty_array", serde_json::json!([])),
("null", serde_json::Value::Null),
(
"unicode_and_escapes",
serde_json::json!({
"s": "héllo \"quoted\" \n \t \u{0} emoji \u{1F600} \u{2028}"
}),
),
(
"numbers",
serde_json::json!({
"u64_max": u64::MAX,
"i64_min": i64::MIN,
"zero": 0,
"neg_zero_float": -0.0_f64,
"integral_float": 1.0_f64,
"fractional": 1234.567890123_f64,
"small_exp": 1.5e-10_f64,
"large_exp": 1.5e300_f64,
}),
),
(
"nested",
serde_json::json!({
"a": [1, 2, {"b": [null, true, false, "x"]}],
"c": {}
}),
),
];
for (name, value) in cases {
let serde_bytes = serde_json::to_vec(&value).expect("serde_json must serialize");
let sonic_bytes = sonic_rs::to_vec(&value).expect("sonic_rs must serialize");
let serde_roundtrip: serde_json::Value =
serde_json::from_slice(&serde_bytes).expect("serde_json bytes must parse");
let sonic_roundtrip: serde_json::Value =
serde_json::from_slice(&sonic_bytes).expect("sonic_rs bytes must parse");
assert_eq!(
serde_roundtrip, sonic_roundtrip,
"case `{name}`: sonic_rs and serde_json must be semantically equivalent"
);
if serde_bytes != sonic_bytes {
eprintln!(
"note: case `{name}` byte output differs (semantically equal) — \
serde_json={:?} sonic_rs={:?}",
String::from_utf8_lossy(&serde_bytes),
String::from_utf8_lossy(&sonic_bytes)
);
}
}
}
}