use std::sync::Arc;
use bytes::Bytes;
use fraiseql_core::{db::types::ColumnSpec, schema::TypeDefinition, security::SecurityContext};
use futures::{StreamExt as _, stream};
use http_body::Frame;
use prost::Message as _;
use prost_reflect::MessageDescriptor;
use tracing::debug;
use super::handler;
fn grpc_frame(msg_bytes: &[u8]) -> Bytes {
let len = msg_bytes.len();
let mut framed = Vec::with_capacity(5 + len);
framed.push(0); #[allow(clippy::cast_possible_truncation)]
framed.extend_from_slice(&(len as u32).to_be_bytes());
framed.extend_from_slice(msg_bytes);
Bytes::from(framed)
}
struct StreamState {
chunks: futures::stream::ReadyChunks<fraiseql_core::db::traits::ColumnRowStream>,
columns: Vec<ColumnSpec>,
row_descriptor: MessageDescriptor,
sent_trailers: bool,
}
fn error_body(
message: String,
) -> impl futures::Stream<Item = Result<Frame<Bytes>, std::convert::Infallible>> + Send {
stream::once(async move { Ok(error_trailers(&message)) })
}
#[allow(clippy::too_many_arguments)]
pub async fn build_streaming_body(
executor: Arc<fraiseql_core::runtime::Executor>,
query_name: String,
columns: Vec<ColumnSpec>,
row_descriptor: MessageDescriptor,
type_def: &TypeDefinition,
request_msg: &prost_reflect::DynamicMessage,
security_context: Option<&SecurityContext>,
batch_size: u32,
) -> impl futures::Stream<Item = Result<Frame<Bytes>, std::convert::Infallible>> + Send {
let query_match = match handler::grpc_query_match(
executor.schema(),
&query_name,
&columns,
true,
request_msg,
type_def,
) {
Ok(qm) => qm,
Err(e) => return futures::future::Either::Left(error_body(e.to_string())),
};
debug!(query = %query_name, batch_size, "Opening gRPC streaming read through the engine");
let opened = executor.stream_row_read(&query_match, None, security_context, &columns).await;
let (columns, rows) = match opened {
Ok(read) => (read.columns, read.stream),
Err(e) => {
return futures::future::Either::Left(error_body(e.to_string()));
},
};
let framed = stream::unfold(
StreamState {
chunks: rows.ready_chunks(usize::try_from(batch_size.max(1)).unwrap_or(usize::MAX)),
columns,
row_descriptor,
sent_trailers: false,
},
|mut state| async move {
if state.sent_trailers {
return None;
}
let Some(chunk) = state.chunks.next().await else {
state.sent_trailers = true;
return Some((Ok(Frame::trailers(ok_trailers())), state));
};
let mut all_frames = Vec::new();
for row in chunk {
match row {
Ok(row) => {
let row_msg =
handler::encode_row(&row, &state.columns, &state.row_descriptor);
all_frames.extend_from_slice(&grpc_frame(&row_msg.encode_to_vec()));
},
Err(e) => {
state.sent_trailers = true;
return Some((Ok(error_trailers(&e.to_string())), state));
},
}
}
Some((Ok(Frame::data(Bytes::from(all_frames))), state))
},
);
futures::future::Either::Right(framed)
}
fn ok_trailers() -> http::HeaderMap {
let mut trailers = http::HeaderMap::new();
trailers.insert("grpc-status", http::HeaderValue::from_static("0"));
trailers
}
fn error_trailers(message: &str) -> Frame<Bytes> {
let mut trailers = http::HeaderMap::new();
trailers.insert("grpc-status", http::HeaderValue::from_static("13"));
if let Ok(msg) = http::HeaderValue::from_str(message) {
trailers.insert("grpc-message", msg);
}
Frame::trailers(trailers)
}