use crate::query_planner::planner::plan_nodes::CustomScalarPaths;
use bytes::{Buf, Bytes};
use futures::stream::BoxStream;
use http_body_util::BodyExt;
use hyper::body::Body;
use crate::executor::{
executors::error::SubgraphExecutorError, response::subgraph_response::SubgraphResponse,
};
const MAX_BUFFER_SIZE: usize = 10 * 1024 * 1024;
#[derive(thiserror::Error, Debug)]
pub enum ParseError {
#[error("Invalid UTF-8 sequence: {0}")]
InvalidUtf8(String),
#[error("Stream read error: {0}")]
StreamReadError(String),
#[error("Invalid subgraph response: {0}")]
InvalidSubgraphResponse(SubgraphExecutorError),
#[error("Buffer size limit exceeded: stream sent more than {MAX_BUFFER_SIZE} bytes without an event boundary")]
BufferSizeLimitExceeded,
}
pub fn parse_to_stream<B>(
body_stream: B,
custom_scalar_paths: Option<CustomScalarPaths>,
) -> BoxStream<'static, Result<SubgraphResponse<'static>, ParseError>>
where
B: Body + Send + Unpin + 'static,
B::Data: Buf + Send,
B::Error: std::fmt::Display + Send,
{
let stream = async_stream::stream! {
let mut body = body_stream;
let mut buffer = Vec::<u8>::new();
loop {
while let Some(boundary) = find_sse_event_boundary(&buffer) {
let event_bytes: Vec<u8> = buffer.drain(..boundary).collect();
match parse(&event_bytes) {
Ok(Some(sse_event)) => {
match sse_event.event.as_deref() {
Some("next") if !sse_event.data.is_empty() => {
match SubgraphResponse::deserialize_from_bytes(
Bytes::from(sse_event.data.clone()),
custom_scalar_paths.as_ref(),
) {
Ok(response) => {
yield Ok(response);
}
Err(e) => {
yield Err(ParseError::InvalidSubgraphResponse(e));
return;
}
}
}
Some("complete") => {
return;
}
_ => {
}
}
}
Err(e) => {
yield Err(e);
return;
}
_ => {}
}
}
match body.frame().await {
Some(Ok(frame)) => {
if let Ok(data) = frame.into_data() {
buffer.extend_from_slice(data.chunk());
if buffer.len() > MAX_BUFFER_SIZE {
yield Err(ParseError::BufferSizeLimitExceeded);
return;
}
}
}
Some(Err(e)) => {
yield Err(ParseError::StreamReadError(e.to_string()));
return;
}
None => {
return;
}
}
}
};
Box::pin(stream)
}
#[derive(Debug)]
struct SubgraphSseEvent {
pub event: Option<String>,
pub data: String,
}
fn find_sse_event_boundary(buffer: &[u8]) -> Option<usize> {
for i in 0..buffer.len().saturating_sub(1) {
if buffer[i] == b'\r'
&& i + 3 < buffer.len()
&& buffer[i + 1] == b'\n'
&& buffer[i + 2] == b'\r'
&& buffer[i + 3] == b'\n'
{
return Some(i + 4);
}
if buffer[i] == b'\n' && buffer[i + 1] == b'\n' {
return Some(i + 2);
}
}
None
}
fn parse(raw: &[u8]) -> Result<Option<SubgraphSseEvent>, ParseError> {
let text = std::str::from_utf8(raw).map_err(|e| ParseError::InvalidUtf8(e.to_string()))?;
let mut current_event: Option<String> = None;
let mut current_data_lines: Vec<String> = Vec::new();
for line in text.lines() {
if line.is_empty() {
if current_event.is_some() || !current_data_lines.is_empty() {
return Ok(Some(SubgraphSseEvent {
event: current_event,
data: current_data_lines.join("\n"),
}));
}
continue;
}
if line.starts_with(':') {
continue;
}
if let Some(colon_pos) = line.find(':') {
let field = &line[..colon_pos];
let value = &line[colon_pos + 1..];
let value = value.trim();
match field {
"event" => {
current_event = Some(value.to_string());
}
"data" => {
current_data_lines.push(value.to_string());
}
_ => {
}
}
}
}
if current_event.is_some() || !current_data_lines.is_empty() {
return Ok(Some(SubgraphSseEvent {
event: current_event,
data: current_data_lines.join("\n"),
}));
}
Ok(None)
}
#[cfg(test)]
mod tests {
use super::*;
use crate::query_planner::planner::plan_nodes::CustomScalarPaths;
#[test]
fn test_parse_single_event_with_data() {
let sse_data = br#"event: next
data: some data
"#;
let event = parse(sse_data).expect("Should parse valid SSE");
assert!(event.is_some());
let event = event.unwrap();
assert_eq!(event.event, Some("next".to_string()));
assert_eq!(event.data, "some data");
}
#[test]
fn test_parse_event_without_explicit_type() {
let sse_data = b"data: some data\n\n";
let event = parse(sse_data).expect("Should parse valid SSE");
assert!(event.is_some());
let event = event.unwrap();
assert_eq!(event.event, None);
assert_eq!(event.data, "some data");
}
#[test]
fn test_parse_just_event() {
let sse_data = b"event: complete\n\n";
let event = parse(sse_data).expect("Should parse valid SSE");
assert!(event.is_some());
let event = event.unwrap();
assert_eq!(event.event, Some("complete".to_string()));
assert_eq!(event.data, "");
}
#[test]
fn test_parse_multiple_events() {
let sse_data = br#"event: next
data: value 1
event: next
data: value 2
event: complete
"#;
let event = parse(sse_data).expect("Should parse valid SSE");
assert!(event.is_some());
let event = event.unwrap();
assert_eq!(event.event, Some("next".to_string()));
assert_eq!(event.data, "value 1");
}
#[test]
fn test_parse_multiline_data() {
let sse_data = b"event: next\ndata: line1\ndata: line2\n\n";
let event = parse(sse_data).expect("Should parse valid SSE");
assert!(event.is_some());
let event = event.unwrap();
assert_eq!(event.event, Some("next".to_string()));
assert_eq!(event.data, "line1\nline2");
}
#[test]
fn test_parse_no_double_newline() {
let sse_data = b"event: next\ndata: line0";
let event = parse(sse_data).expect("Should parse valid SSE");
assert!(event.is_some());
let event = event.unwrap();
assert_eq!(event.event, Some("next".to_string()));
assert_eq!(event.data, "line0");
}
#[test]
fn test_parse_heartbeat() {
let sse_data = b":\n\nevent: next\ndata: payload\n\n:\n\n";
let event = parse(sse_data).expect("Should parse valid SSE");
assert!(event.is_some());
let event = event.unwrap();
assert_eq!(event.event, Some("next".to_string()));
assert_eq!(event.data, "payload");
}
#[test]
fn test_parse_empty_input() {
let sse_data = b"";
let event = parse(sse_data).expect("Should handle empty input");
assert!(event.is_none());
}
#[test]
fn test_find_sse_event_boundary_crlf() {
let buffer = b"event: next\r\ndata: hello\r\n\r\n";
let boundary = find_sse_event_boundary(buffer);
assert_eq!(boundary, Some(buffer.len()));
}
#[test]
fn test_find_sse_event_boundary_lf() {
let buffer = b"event: next\ndata: hello\n\n";
let boundary = find_sse_event_boundary(buffer);
assert_eq!(boundary, Some(buffer.len()));
}
#[test]
fn test_parse_single_event_crlf() {
let sse_data = b"event: next\r\ndata: some data\r\n\r\n";
let event = parse(sse_data).expect("Should parse valid SSE with CRLF");
assert!(event.is_some());
let event = event.unwrap();
assert_eq!(event.event, Some("next".to_string()));
assert_eq!(event.data, "some data");
}
#[test]
fn test_parse_just_event_crlf() {
let sse_data = b"event: complete\r\n\r\n";
let event = parse(sse_data).expect("Should parse valid SSE with CRLF");
assert!(event.is_some());
let event = event.unwrap();
assert_eq!(event.event, Some("complete".to_string()));
assert_eq!(event.data, "");
}
#[tokio::test]
async fn test_parse_to_stream_chunked_events_crlf() {
use bytes::Bytes;
use futures::StreamExt;
use http_body_util::StreamBody;
use hyper::body::Frame;
let chunks: Vec<Result<Frame<Bytes>, std::convert::Infallible>> = vec![
Ok(Frame::data(Bytes::from(
"event: next\r\ndata: {\"data\":{\"hello\":\"wor",
))),
Ok(Frame::data(Bytes::from("ld\"}}\r\n\r\neve"))),
Ok(Frame::data(Bytes::from("nt: complete\r\n\r\n"))),
];
let body = StreamBody::new(futures::stream::iter(chunks));
let mut stream = parse_to_stream(body, None);
let first = stream.next().await;
assert!(first.is_some());
let first_result = first.unwrap();
assert!(first_result.is_ok());
let response = first_result.unwrap();
assert!(!response.data.is_null());
let second = stream.next().await;
assert!(second.is_none());
}
#[tokio::test]
async fn test_parse_to_stream_chunked_events() {
use bytes::Bytes;
use futures::StreamExt;
use http_body_util::StreamBody;
use hyper::body::Frame;
let chunks: Vec<Result<Frame<Bytes>, std::convert::Infallible>> = vec![
Ok(Frame::data(Bytes::from(
"event: next\ndata: {\"data\":{\"hello\":\"wor",
))),
Ok(Frame::data(Bytes::from("ld\"}}\n\neve"))),
Ok(Frame::data(Bytes::from("nt: complete\n\n"))),
];
let body = StreamBody::new(futures::stream::iter(chunks));
let mut stream = parse_to_stream(body, None);
let first = stream.next().await;
assert!(first.is_some());
let first_result = first.unwrap();
assert!(first_result.is_ok());
let response = first_result.unwrap();
assert!(!response.data.is_null());
let second = stream.next().await;
assert!(second.is_none());
}
#[tokio::test]
async fn test_parse_to_stream_uses_custom_scalar_paths() {
use bytes::Bytes;
use futures::StreamExt;
use http_body_util::StreamBody;
use hyper::body::Frame;
let chunks: Vec<Result<Frame<Bytes>, std::convert::Infallible>> = vec![Ok(Frame::data(
Bytes::from(
"event: next\ndata: {\"data\":{\"custom\":{\"escaped.key\\t\":\"value\"}}}\n\nevent: complete\n\n",
),
))];
let mut custom_scalar_paths = CustomScalarPaths::default();
custom_scalar_paths.insert_path(["custom"]);
let body = StreamBody::new(futures::stream::iter(chunks));
let mut stream = parse_to_stream(body, Some(custom_scalar_paths));
let first = stream.next().await.unwrap().unwrap();
let data = first.data.as_object().unwrap();
assert!(data[0].1.as_raw_json().is_some());
}
}