use std::collections::HashMap;
use axum::response::Response;
use fraiseql_core::{
graphql::{
defer, parse_query_with_operation_name, selection_set, selection_set::variables_map,
stream_split, types::FieldSelection, value_json,
},
runtime::{QueryMatcher, coerce_pagination_arg},
security::SecurityContext,
};
use futures::{StreamExt as _, stream};
use serde_json::{Value, json};
use tracing::{debug, warn};
use super::{
super::{app_state::AppState, request::GraphQLRequest, tenant_dispatch},
execute_graphql_request,
incremental::{self, Chunk, Wire},
stages,
};
use crate::{
error::{ErrorResponse, GraphQLError},
middleware::stream_auth::StreamAuthGuard,
};
struct StreamPlan {
response_key: String,
initial_count: u64,
client_limit: Option<u64>,
client_offset: u64,
}
pub(in super::super) async fn handle_sse(
state: AppState,
wire: Wire,
headers: axum::http::HeaderMap,
peer_ip: String,
security_context: Option<SecurityContext>,
token_claims: Option<crate::middleware::oidc_auth::SessionTokenClaims>,
mut request: GraphQLRequest,
) -> Result<Response, ErrorResponse> {
let security_context =
Box::pin(stages::authenticate(&state, &headers, security_context)).await?;
let query = Box::pin(stages::resolve_query_body(&state, &mut request)).await?;
let operation_name = request.operation_name.clone();
let op = operation_name.as_deref();
let plan = plan_stream(&state, &query, request.variables.as_ref(), op)?;
let max_depth = state.executor.load().max_query_depth();
let defer_plan = plan_defer(&query, request.variables.as_ref(), op, max_depth);
let nested_stream_plan = plan_nested_stream(&query, request.variables.as_ref(), op, max_depth);
if plan.is_some() && defer_plan.is_some() {
return Err(ErrorResponse::from_error(GraphQLError::new(
"@defer cannot be combined with @stream in one operation: the two order the \
same response differently and interleaving them is not defined here"
.to_string(),
crate::error::ErrorCode::ValidationError,
)));
}
if nested_stream_plan.is_some() && defer_plan.is_some() {
return Err(ErrorResponse::from_error(GraphQLError::new(
"@defer cannot be combined with a nested @stream in one operation: both split \
the delivery of one result and their payload order is not defined here"
.to_string(),
crate::error::ErrorCode::ValidationError,
)));
}
if plan.is_some() && nested_stream_plan.is_some() {
return Err(ErrorResponse::from_error(GraphQLError::new(
"a root @stream cannot be combined with a nested @stream in one operation: the \
root one pages the database and each of its batches would carry its own copy \
of the nested list, which has no incremental addressing"
.to_string(),
crate::error::ErrorCode::ValidationError,
)));
}
let Some(plan) = plan else {
let variables = request.variables.clone();
let batch_size = state.graphql_incremental_batch_size.max(1) as usize;
let response = Box::pin(execute_graphql_request(
state,
request,
None,
security_context,
&headers,
&peer_ip,
))
.await?;
let chunks = match (defer_plan, nested_stream_plan) {
(Some(selections), _) => {
deferred_chunks(response.body, &selections, variables.as_ref())
},
(None, Some(selections)) => {
streamed_chunks(response.body, &selections, variables.as_ref(), batch_size)?
},
(None, None) => vec![Chunk {
payload: response.body,
resume_id: None,
}],
};
return Ok(incremental::respond(wire, stream::iter(chunks)));
};
stages::enforce_introspection_policy(&state, &query, security_context.as_ref())?;
stages::validate_request(&state, &query, &request, &peer_ip)?;
let tenant_key =
tenant_dispatch::resolve_tenant_key(&state, security_context.as_ref(), &headers).map_err(
|e| {
ErrorResponse::from_error(GraphQLError::new(
e.to_string(),
crate::error::ErrorCode::ValidationError,
))
},
)?;
state.metrics.queries_total.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
{
let dispatch = tenant_dispatch::dispatch_to_tenant(&state, tenant_key.as_deref())
.map_err(|e| ErrorResponse::from_error(super::tenant_dispatch_error(&e)))?;
let estimated_cost = tenant_dispatch::estimate_request_cost(
&query,
request.variables.as_ref(),
&dispatch.executor,
);
tenant_dispatch::charge_cost_budget(
&state,
tenant_key.as_deref(),
security_context.as_ref(),
estimated_cost,
)
.map_err(|e| ErrorResponse::from_error(super::tenant_dispatch_error(&e)))?;
}
let base_variables = match request.variables.clone() {
Some(Value::Object(map)) => map,
_ => serde_json::Map::new(),
};
let resume_from = resume_offset(&headers, &plan).map_err(|msg| {
ErrorResponse::from_error(GraphQLError::new(msg, crate::error::ErrorCode::ValidationError))
})?;
let start_offset = resume_from.unwrap_or(plan.client_offset);
let already_delivered = start_offset - plan.client_offset;
let remaining_budget = plan.client_limit.map(|total| total.saturating_sub(already_delivered));
let initial_requested =
remaining_budget.map_or(plan.initial_count, |left| plan.initial_count.min(left));
let initial_vars = batch_variables(&base_variables, initial_requested, start_offset);
let first = run_batch(
&state,
tenant_key.as_deref(),
&query,
&initial_vars,
security_context.as_ref(),
op,
)
.await
.map_err(ErrorResponse::from_error)?;
let initial_rows = extract_items(&first, &plan.response_key).map_or(0, <[Value]>::len) as u64;
let delivered = already_delivered + initial_rows;
let exhausted = initial_rows != initial_requested
|| plan.client_limit.is_some_and(|total| delivered >= total);
let mut initial_payload = first;
attach_has_next(&mut initial_payload, !exhausted);
debug!(
response_key = %plan.response_key,
initial_rows,
resume_from = ?resume_from,
exhausted,
"GraphQL @stream delivery started"
);
let auth_guard = StreamAuthGuard::new(
security_context.as_ref(),
token_claims,
state.revocation_manager.clone(),
);
let batch_size = u64::from(state.graphql_incremental_batch_size.max(1));
let unfold_state = BatchState {
state: state.clone(),
query,
operation_name,
base_variables,
security_context,
auth_guard,
tenant_key,
response_key: plan.response_key,
client_limit: plan.client_limit,
batch_size,
offset: start_offset + initial_rows,
phase: if exhausted {
Phase::Complete
} else {
Phase::Streaming
},
emitted_rows: delivered,
batches: 1,
};
let resume_id = start_offset + initial_rows;
let chunks = stream::iter(vec![Chunk {
payload: initial_payload,
resume_id: Some(resume_id),
}])
.chain(stream::unfold(unfold_state, batch_step));
Ok(incremental::respond(wire, chunks))
}
fn plan_defer(
query: &str,
variables: Option<&Value>,
operation_name: Option<&str>,
max_depth: u32,
) -> Option<Vec<FieldSelection>> {
let parsed = parse_query_with_operation_name(query, operation_name).ok()?;
let vars = variables_map(variables);
let effective =
selection_set::resolve_and_filter(&parsed.selections, &parsed.fragments, &vars, max_depth)
.ok()?;
defer::contains_defer(&effective, &vars).then_some(effective)
}
fn plan_nested_stream(
query: &str,
variables: Option<&Value>,
operation_name: Option<&str>,
max_depth: u32,
) -> Option<Vec<FieldSelection>> {
let parsed = parse_query_with_operation_name(query, operation_name).ok()?;
let vars = variables_map(variables);
let effective =
selection_set::resolve_and_filter(&parsed.selections, &parsed.fragments, &vars, max_depth)
.ok()?;
stream_split::contains_nested_stream(&effective, &vars).then_some(effective)
}
fn streamed_chunks(
mut body: Value,
selections: &[FieldSelection],
variables: Option<&Value>,
batch_size: usize,
) -> Result<Vec<Chunk>, ErrorResponse> {
let vars = variables_map(variables);
let streamed = match body.get_mut("data") {
Some(data) => stream_split::split(selections, data, &vars, batch_size).map_err(|e| {
ErrorResponse::from_error(GraphQLError::new(
format!("@stream on `{}`: {}", e.field, e.reason),
crate::error::ErrorCode::ValidationError,
))
})?,
None => Vec::new(),
};
if streamed.is_empty() {
return Ok(vec![Chunk {
payload: body,
resume_id: None,
}]);
}
attach_has_next(&mut body, true);
let mut chunks = vec![Chunk {
payload: body,
resume_id: None,
}];
let last = streamed.len() - 1;
for (index, chunk) in streamed.into_iter().enumerate() {
chunks.push(Chunk {
payload: json!({
"incremental": [stream_split::incremental_entry(chunk)],
"hasNext": index != last,
}),
resume_id: None,
});
}
Ok(chunks)
}
fn deferred_chunks(
mut body: Value,
selections: &[FieldSelection],
variables: Option<&Value>,
) -> Vec<Chunk> {
let vars = variables_map(variables);
let deferred = body
.get_mut("data")
.map(|data| defer::split(selections, data, &vars))
.unwrap_or_default();
if deferred.is_empty() {
return vec![Chunk {
payload: body,
resume_id: None,
}];
}
attach_has_next(&mut body, true);
let mut chunks = vec![Chunk {
payload: body,
resume_id: None,
}];
let last = deferred.len() - 1;
for (index, payload) in deferred.into_iter().enumerate() {
let mut entry = json!({
"data": Value::Object(payload.data),
"path": payload.path,
});
if let Some(label) = payload.label {
entry["label"] = Value::String(label);
}
chunks.push(Chunk {
payload: json!({
"incremental": [entry],
"hasNext": index != last,
}),
resume_id: None,
});
}
chunks
}
fn resume_offset(
headers: &axum::http::HeaderMap,
plan: &StreamPlan,
) -> Result<Option<u64>, String> {
let Some(raw) = headers.get("last-event-id").and_then(|v| v.to_str().ok()) else {
return Ok(None);
};
let raw = raw.trim();
if raw.is_empty() {
return Ok(None);
}
let offset: u64 = raw.parse().map_err(|_| {
format!(
"Last-Event-ID must be the absolute row offset this transport emits as the \
event id, got {raw:?}"
)
})?;
if offset < plan.client_offset {
return Err(format!(
"Last-Event-ID {offset} precedes the query's own offset argument \
({}); resuming cannot deliver rows before the requested start",
plan.client_offset
));
}
Ok(Some(offset))
}
enum Phase {
Streaming,
Complete,
Finished,
}
struct BatchState {
state: AppState,
query: String,
operation_name: Option<String>,
base_variables: serde_json::Map<String, Value>,
security_context: Option<SecurityContext>,
auth_guard: StreamAuthGuard,
tenant_key: Option<String>,
response_key: String,
client_limit: Option<u64>,
batch_size: u64,
offset: u64,
phase: Phase,
emitted_rows: u64,
batches: u64,
}
async fn batch_step(mut st: BatchState) -> Option<(Chunk, BatchState)> {
match st.phase {
Phase::Finished => None,
Phase::Complete => {
st.phase = Phase::Finished;
st.state
.metrics
.queries_success
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
tracing::info!(
target: "fraiseql::sse_audit",
batches = st.batches,
rows = st.emitted_rows,
tenant = st.tenant_key.as_deref().unwrap_or(""),
"@stream delivery complete"
);
None
},
Phase::Streaming => {
if let Err(reason) = st.auth_guard.check().await {
warn!(reason, "@stream delivery terminated: principal no longer valid");
st.phase = Phase::Complete;
let payload = json!({
"errors": [{
"message": format!("{reason} during streaming delivery"),
"extensions": {"code": "UNAUTHENTICATED"}
}],
"hasNext": false,
});
return Some((resumable_chunk(payload, st.offset), st));
}
let remaining = st.client_limit.map(|total| total.saturating_sub(st.emitted_rows));
let requested = remaining.map_or(st.batch_size, |r| r.min(st.batch_size));
if requested == 0 {
st.phase = Phase::Complete;
return Some((resumable_chunk(json!({"hasNext": false}), st.offset), st));
}
let vars = batch_variables(&st.base_variables, requested, st.offset);
let result = run_batch(
&st.state,
st.tenant_key.as_deref(),
&st.query,
&vars,
st.security_context.as_ref(),
st.operation_name.as_deref(),
)
.await;
st.batches += 1;
match result {
Err(err) => {
st.phase = Phase::Complete;
st.state
.metrics
.queries_error
.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let err_json = serde_json::to_value(&err)
.unwrap_or_else(|_| json!({"message": "streamed batch failed"}));
let payload = json!({
"errors": [err_json],
"hasNext": false,
});
Some((resumable_chunk(payload, st.offset), st))
},
Ok(response) => {
let items = extract_items(&response, &st.response_key)
.map(<[Value]>::to_vec)
.unwrap_or_default();
let count = items.len() as u64;
let path_index = st.offset;
st.offset += count;
st.emitted_rows += count;
if count > requested {
warn!(
count,
requested,
"@stream batch returned more rows than requested; \
pagination is not binding — terminating the delivery"
);
}
let done = count != requested
|| st.client_limit.is_some_and(|total| st.emitted_rows >= total);
if done {
st.phase = Phase::Complete;
}
let payload = json!({
"incremental": [{
"items": items,
"path": [st.response_key, path_index],
}],
"hasNext": !done,
});
let resume_id = st.offset;
Some((resumable_chunk(payload, resume_id), st))
},
}
},
}
}
#[allow(clippy::result_large_err)]
async fn run_batch(
state: &AppState,
tenant_key: Option<&str>,
query: &str,
variables: &Value,
security_context: Option<&SecurityContext>,
operation_name: Option<&str>,
) -> Result<Value, GraphQLError> {
let dispatch = tenant_dispatch::dispatch_to_tenant(state, tenant_key)
.map_err(|e| super::tenant_dispatch_error(&e))?;
let executor = &dispatch.executor;
let result = if let Some(ctx) = security_context {
executor
.execute_operation_with_security(query, Some(variables), ctx, operation_name)
.await
} else {
executor.execute_operation(query, Some(variables), operation_name).await
};
#[allow(unused_mut)]
let mut response = result
.map_err(|e| state.error_sanitizer.sanitize(GraphQLError::from_fraiseql_error(&e)))?;
#[cfg(feature = "secrets")]
stages::decrypt_response_fields(state, &mut response)
.await
.map_err(|_| GraphQLError::internal("response post-processing failed"))?;
Ok(response)
}
fn batch_variables(base: &serde_json::Map<String, Value>, limit: u64, offset: u64) -> Value {
let mut vars = base.clone();
vars.insert("limit".to_string(), json!(limit));
vars.insert("offset".to_string(), json!(offset));
Value::Object(vars)
}
fn extract_items<'a>(response: &'a Value, response_key: &str) -> Option<&'a [Value]> {
response.get("data")?.get(response_key)?.as_array().map(Vec::as_slice)
}
fn attach_has_next(payload: &mut Value, has_next: bool) {
if let Some(obj) = payload.as_object_mut() {
obj.insert("hasNext".to_string(), Value::Bool(has_next));
}
}
const fn resumable_chunk(payload: Value, next_offset: u64) -> Chunk {
Chunk {
payload,
resume_id: Some(next_offset),
}
}
fn plan_stream(
state: &AppState,
query: &str,
variables: Option<&Value>,
operation_name: Option<&str>,
) -> Result<Option<StreamPlan>, ErrorResponse> {
let bad_request = |msg: &str| {
ErrorResponse::from_error(GraphQLError::new(
msg.to_string(),
crate::error::ErrorCode::ValidationError,
))
};
let Ok(parsed) = parse_query_with_operation_name(query, operation_name) else {
return Ok(None);
};
if !selections_contain_stream(&parsed.selections) {
return Ok(None);
}
if parsed.operation_type != "query" {
return Err(bad_request("@stream is only supported on query operations"));
}
if !parsed
.selections
.iter()
.any(|s| s.directives.iter().any(|d| d.name == "stream"))
{
return Ok(None);
}
if parsed.variables.iter().any(|v| v.name == "limit" || v.name == "offset") {
return Err(bad_request(
"@stream cannot be combined with document variables named $limit or $offset: \
streaming paginates by injecting those variables. Rename the variables.",
));
}
let defaults = fraiseql_core::graphql::value_json::variable_defaults(&parsed.variables)
.map_err(|e| bad_request(&format!("@stream planning failed: {e}")))?;
let defaulted =
fraiseql_core::graphql::value_json::with_variable_defaults(&defaults, variables);
let variables = defaulted.as_ref().or(variables);
let executor = state.executor.load();
let matcher = QueryMatcher::new(executor.schema().clone());
let matched = matcher
.match_query(query, variables)
.map_err(|e| bad_request(&format!("@stream planning failed: {e}")))?;
if matched.fields.len() != 1 {
return Err(bad_request("@stream requires exactly one root field in the operation"));
}
let root = matched
.selections
.first()
.ok_or_else(|| bad_request("@stream requires a selected root field"))?;
let Some(directive) = root.directives.iter().find(|d| d.name == "stream") else {
return Ok(None);
};
let vars_map = variables_map(variables);
if !stream_enabled(directive, &vars_map).map_err(|msg| bad_request(&msg))? {
return Ok(None);
}
let initial_count = directive_u64_arg(directive, "initialCount", &vars_map)
.map_err(|msg| bad_request(&msg))?
.unwrap_or(0);
let def = &matched.query_def;
if !def.returns_list {
return Err(bad_request("@stream requires a list-returning query field"));
}
if def.relay {
return Err(bad_request(
"@stream is not supported on relay (connection) queries; use cursor \
pagination instead",
));
}
if !(def.auto_params.has_limit && def.auto_params.has_offset) {
return Err(bad_request(
"@stream requires the query to accept limit and offset parameters",
));
}
let client_limit = coerce_pagination_arg("limit", matched.arguments.get("limit"))
.map_err(|e| bad_request(&format!("@stream planning failed: {e}")))?
.map(u64::from);
let client_offset = coerce_pagination_arg("offset", matched.arguments.get("offset"))
.map_err(|e| bad_request(&format!("@stream planning failed: {e}")))?
.map_or(0, u64::from);
Ok(Some(StreamPlan {
response_key: matched.response_key().to_string(),
initial_count,
client_limit,
client_offset,
}))
}
fn selections_contain_stream(selections: &[FieldSelection]) -> bool {
selections.iter().any(|s| {
s.directives.iter().any(|d| d.name == "stream")
|| selections_contain_stream(&s.nested_fields)
})
}
fn stream_enabled(
directive: &fraiseql_core::graphql::types::Directive,
variables: &HashMap<String, Value>,
) -> Result<bool, String> {
match resolve_directive_arg(directive, "if", variables)? {
None => Ok(true),
Some(Value::Bool(b)) => Ok(b),
Some(other) => Err(format!("@stream(if:) must be a Boolean, got {other}")),
}
}
fn directive_u64_arg(
directive: &fraiseql_core::graphql::types::Directive,
name: &str,
variables: &HashMap<String, Value>,
) -> Result<Option<u64>, String> {
match resolve_directive_arg(directive, name, variables)? {
None => Ok(None),
Some(value) => value
.as_u64()
.map(Some)
.ok_or_else(|| format!("@stream({name}:) must be a non-negative integer, got {value}")),
}
}
fn resolve_directive_arg(
directive: &fraiseql_core::graphql::types::Directive,
name: &str,
variables: &HashMap<String, Value>,
) -> Result<Option<Value>, String> {
let Some(arg) = directive.arguments.iter().find(|a| a.name == name) else {
return Ok(None);
};
let decoded = value_json::decode(&arg.value_json)
.map_err(|e| format!("@stream({name}:) could not be decoded: {e}"))?;
Ok(Some(value_json::resolve_variables(decoded, variables)))
}