use std::{sync::atomic::Ordering, time::Instant};
use axum::{
Json,
extract::{Query, State},
http::HeaderMap,
};
use fraiseql_core::{
apq::{ApqMetrics, ApqStorage},
security::SecurityContext,
};
use fraiseql_error::FraiseQLError;
use tracing::{debug, error, warn};
use super::{
app_state::AppState,
request::{GraphQLGetParams, GraphQLRequest, GraphQLResponse},
};
use crate::{
error::{ErrorResponse, GraphQLError},
extractors::{OptionalSecurityContext, PeerIp},
tracing_utils,
};
#[tracing::instrument(skip_all, fields(operation_name))]
#[doc(hidden)] pub async fn graphql_handler(
State(state): State<AppState>,
headers: HeaderMap,
PeerIp(peer_ip): PeerIp,
OptionalSecurityContext(security_context): OptionalSecurityContext,
token_claims: Option<axum::Extension<crate::middleware::oidc_auth::SessionTokenClaims>>,
Json(request): Json<GraphQLRequest>,
) -> Result<axum::response::Response, ErrorResponse> {
let trace_context = tracing_utils::extract_trace_context(&headers);
if trace_context.is_some() {
debug!("Extracted W3C trace context from incoming request");
}
if security_context.is_some() {
debug!("Authenticated request with security context");
}
if let Some(wire) = state
.graphql_incremental_enabled
.then(|| incremental::negotiate(&headers))
.flatten()
{
return Box::pin(sse::handle_sse(
state,
wire,
headers,
peer_ip,
security_context,
token_claims.map(|axum::Extension(claims)| claims),
request,
))
.await;
}
execute_graphql_request(state, request, trace_context, security_context, &headers, &peer_ip)
.await
.map(axum::response::IntoResponse::into_response)
}
#[tracing::instrument(skip_all, fields(operation_name))]
#[doc(hidden)] pub async fn graphql_get_handler(
State(state): State<AppState>,
headers: HeaderMap,
PeerIp(peer_ip): PeerIp,
OptionalSecurityContext(security_context): OptionalSecurityContext,
token_claims: Option<axum::Extension<crate::middleware::oidc_auth::SessionTokenClaims>>,
Query(params): Query<GraphQLGetParams>,
) -> Result<axum::response::Response, ErrorResponse> {
let max_get_bytes = state.max_get_query_bytes;
if params.query.len() > max_get_bytes {
return Err(ErrorResponse::from_error(GraphQLError::payload_too_large(format!(
"GET query string exceeds maximum allowed length ({max_get_bytes} bytes)"
))));
}
let variables = if let Some(vars_str) = params.variables {
if vars_str.len() > max_get_bytes {
return Err(ErrorResponse::from_error(GraphQLError::payload_too_large(format!(
"GET variables string exceeds maximum allowed length ({max_get_bytes} bytes)"
))));
}
match serde_json::from_str::<serde_json::Value>(&vars_str) {
Ok(v) => Some(v),
Err(e) => {
warn!(
error = %e,
variables_bytes = vars_str.len(),
"Failed to parse variables JSON in GET request"
);
return Err(ErrorResponse::from_error(GraphQLError::request(format!(
"Invalid variables JSON: {e}"
))));
},
}
} else {
None
};
if detect_mutation_name(¶ms.query).is_some() {
warn!(
operation_name = ?params.operation_name,
"Mutation sent via GET request — rejected (use POST)"
);
return Err(ErrorResponse::from_error(GraphQLError::method_not_allowed(
"Mutations must be sent over POST, not GET",
)));
}
let trace_context = tracing_utils::extract_trace_context(&headers);
if trace_context.is_some() {
debug!("Extracted W3C trace context from incoming request");
}
let request = GraphQLRequest {
query: Some(params.query),
variables,
operation_name: params.operation_name,
extensions: None,
document_id: None,
};
if security_context.is_some() {
debug!("Authenticated GET request with security context");
}
if let Some(wire) = state
.graphql_incremental_enabled
.then(|| incremental::negotiate(&headers))
.flatten()
{
return Box::pin(sse::handle_sse(
state,
wire,
headers,
peer_ip,
security_context,
token_claims.map(|axum::Extension(claims)| claims),
request,
))
.await;
}
execute_graphql_request(state, request, trace_context, security_context, &headers, &peer_ip)
.await
.map(axum::response::IntoResponse::into_response)
}
pub(crate) const HTTP_QUERY_METHOD: &str = "QUERY";
pub async fn graphql_query_method_handler(
State(state): State<AppState>,
method: axum::http::Method,
headers: HeaderMap,
PeerIp(peer_ip): PeerIp,
OptionalSecurityContext(security_context): OptionalSecurityContext,
body: axum::body::Bytes,
) -> Result<GraphQLResponse, ErrorResponse> {
if method.as_str() != HTTP_QUERY_METHOD {
return Err(ErrorResponse::from_error(GraphQLError::method_not_allowed(
"Method not allowed on the GraphQL endpoint",
)));
}
let request: GraphQLRequest = serde_json::from_slice(&body).map_err(|e| {
ErrorResponse::from_error(GraphQLError::request(format!("Invalid QUERY body: {e}")))
})?;
if let Some(op) = non_query_operation(request.query.as_deref().unwrap_or_default()) {
warn!(operation = op, "Non-query operation sent via QUERY — rejected (use POST)");
return Err(ErrorResponse::from_error(GraphQLError::method_not_allowed(
"Only query operations may be sent over the QUERY method; use POST",
)));
}
let trace_context = tracing_utils::extract_trace_context(&headers);
execute_graphql_request(state, request, trace_context, security_context, &headers, &peer_ip)
.await
}
pub(crate) fn non_query_operation(query: &str) -> Option<&'static str> {
let parsed = fraiseql_core::graphql::parse_query(query).ok()?;
match parsed.operation_type.as_str() {
"mutation" => Some("mutation"),
"subscription" => Some("subscription"),
_ => None,
}
}
pub(crate) fn detect_mutation_name(query: &str) -> Option<String> {
let parsed = fraiseql_core::graphql::parse_query(query).ok()?;
if parsed.operation_type == "mutation" {
Some(parsed.root_field)
} else {
None
}
}
#[cfg(feature = "auth")]
#[allow(dead_code)] pub(crate) fn extract_ip_from_headers(_headers: &HeaderMap) -> String {
"unknown".to_string()
}
pub(crate) fn extract_apq_hash(extensions: Option<&serde_json::Value>) -> Option<&str> {
extensions?.get("persistedQuery")?.get("sha256Hash")?.as_str()
}
fn extract_document_id(request: &GraphQLRequest) -> Option<String> {
if let Some(ref doc_id) = request.document_id {
return Some(doc_id.clone());
}
if let Some(ext) = request.extensions.as_ref() {
if let Some(doc_id) = ext.get("doc_id").and_then(|v| v.as_str()) {
return Some(doc_id.to_string());
}
if let Some(hash) = ext
.get("persistedQuery")
.and_then(|pq| pq.get("sha256Hash"))
.and_then(|h| h.as_str())
{
return Some(hash.to_string());
}
}
None
}
pub(crate) async fn resolve_apq(
apq_store: &dyn ApqStorage,
apq_metrics: &ApqMetrics,
hash: &str,
query_body: Option<&str>,
) -> Result<String, ErrorResponse> {
if let Some(body) = query_body {
if !fraiseql_core::apq::verify_hash(body, hash) {
apq_metrics.record_error();
return Err(ErrorResponse::from_error(GraphQLError::persisted_query_mismatch()));
}
if let Err(e) = apq_store.set(hash.to_owned(), body.to_owned()).await {
warn!(error = %e, "Failed to store APQ query — proceeding without caching");
apq_metrics.record_error();
} else {
apq_metrics.record_store();
}
Ok(body.to_owned())
} else {
match apq_store.get(hash).await {
Ok(Some(stored)) => {
apq_metrics.record_hit();
Ok(stored)
},
Ok(None) => {
apq_metrics.record_miss();
Err(ErrorResponse::from_error(GraphQLError::persisted_query_not_found()))
},
Err(e) => {
warn!(error = %e, "APQ store lookup failed — treating as miss");
apq_metrics.record_error();
Err(ErrorResponse::from_error(GraphQLError::persisted_query_not_found()))
},
}
}
}
#[tracing::instrument(skip_all, fields(operation_name = request.operation_name.as_deref().unwrap_or("anonymous")))]
async fn execute_graphql_request(
state: AppState,
mut request: GraphQLRequest,
#[cfg(feature = "federation")] _trace_context: Option<
fraiseql_core::federation::FederationTraceContext,
>,
#[cfg(not(feature = "federation"))] _trace_context: Option<()>,
security_context: Option<SecurityContext>,
headers: &HeaderMap,
peer_ip: &str,
) -> Result<GraphQLResponse, ErrorResponse> {
let mut security_context =
Box::pin(stages::authenticate(&state, headers, security_context)).await?;
security_context = stages::stamp_trace_context(headers, security_context);
#[cfg(feature = "auth")]
Box::pin(stages::enrich_identity(&state, &mut security_context)).await?;
let query = Box::pin(stages::resolve_query_body(&state, &mut request)).await?;
let start_time = Instant::now();
let metrics = &state.metrics;
metrics.queries_total.fetch_add(1, Ordering::Relaxed);
debug!(
query_length = query.len(),
has_variables = request.variables.is_some(),
operation_name = ?request.operation_name,
"Executing GraphQL query"
);
stages::enforce_introspection_policy(&state, &query, security_context.as_ref())?;
stages::validate_request(&state, &query, &request, peer_ip)?;
#[cfg(feature = "federation")]
let cb_entity_types =
stages::check_federation_circuit_breakers(&state, &query, request.variables.as_ref())?;
let tenant_key =
super::tenant_dispatch::resolve_tenant_key(&state, security_context.as_ref(), headers)
.map_err(|e| ErrorResponse::from_error(GraphQLError::from_fraiseql_error(&e)))?;
let idempotency_key = if detect_mutation_name(&query).is_some() {
headers.get("idempotency-key").and_then(|v| v.to_str().ok()).map(|client_key| {
let scope = crate::routes::idempotency::IdempotencyScope {
tenant: tenant_key.clone(),
principal: security_context.as_ref().map(|c| c.user_id.to_string()),
method: "POST".to_string(),
path: "/graphql".to_string(),
};
let body_hash = crate::routes::idempotency::hash_body(&serde_json::json!({
"query": query,
"variables": request.variables,
"operationName": request.operation_name,
}));
(scope.key(client_key), body_hash)
})
} else {
None
};
if let Some((ref key, body_hash)) = idempotency_key {
match state.idempotency_store.check(key, body_hash).await {
crate::routes::idempotency::IdempotencyCheck::Replay(stored) => {
debug!("Replaying stored response for repeated Idempotency-Key mutation");
return Ok(GraphQLResponse {
body: stored.body.unwrap_or(serde_json::Value::Null),
});
},
crate::routes::idempotency::IdempotencyCheck::Conflict => {
return Err(ErrorResponse::from_error(GraphQLError::idempotency_conflict()));
},
crate::routes::idempotency::IdempotencyCheck::New => {},
}
}
let variables = request.variables;
let dispatch = super::tenant_dispatch::dispatch_to_tenant(&state, tenant_key.as_deref())
.map_err(|e| ErrorResponse::from_error(tenant_dispatch_error(&e)))?;
let executor = &dispatch.executor;
let estimated_cost =
super::tenant_dispatch::estimate_request_cost(&query, variables.as_ref(), executor);
super::tenant_dispatch::charge_cost_budget(
&state,
tenant_key.as_deref(),
security_context.as_ref(),
estimated_cost,
)
.map_err(|e| ErrorResponse::from_error(tenant_dispatch_error(&e)))?;
#[cfg(feature = "auth")]
let audit_subject = security_context.as_ref().map(|ctx| ctx.user_id.to_string());
let operation_name = request.operation_name.as_deref();
let exec_result = if let Some(sec_ctx) = security_context {
executor
.execute_operation_with_security(&query, variables.as_ref(), &sec_ctx, operation_name)
.await
} else {
executor.execute_operation(&query, variables.as_ref(), operation_name).await
};
#[cfg(feature = "federation")]
if !cb_entity_types.is_empty() {
if let Some(ref cb_manager) = state.circuit_breaker {
if exec_result.is_ok() {
for entity_type in &cb_entity_types {
cb_manager.record_success(entity_type);
}
} else {
for entity_type in &cb_entity_types {
cb_manager.record_failure(entity_type);
}
}
}
}
let op_name = request.operation_name.as_deref().unwrap_or("");
let result = exec_result.map_err(|e| {
let elapsed = start_time.elapsed();
#[allow(clippy::cast_possible_truncation)]
let elapsed_us = elapsed.as_micros() as u64;
error!(
error = %e,
elapsed_ms = elapsed.as_millis(),
operation_name = ?request.operation_name,
"Query execution failed"
);
metrics.queries_error.fetch_add(1, Ordering::Relaxed);
metrics.execution_errors_total.fetch_add(1, Ordering::Relaxed);
metrics.queries_duration_us.fetch_add(elapsed_us, Ordering::Relaxed);
metrics.operation_metrics.record(op_name, elapsed_us, true);
#[cfg(feature = "auth")]
if matches!(e, fraiseql_core::FraiseQLError::Authorization { .. }) {
use fraiseql_auth::audit::logger::{
AuditEntry, AuditEventType, SecretType, get_audit_logger,
};
let resource =
if let fraiseql_core::FraiseQLError::Authorization { ref resource, .. } = e {
resource.clone().unwrap_or_else(|| op_name.to_string())
} else {
op_name.to_string()
};
get_audit_logger().log_entry(AuditEntry {
event_type: AuditEventType::AuthorizationDenied,
secret_type: SecretType::JwtToken,
subject: audit_subject.clone(),
operation: op_name.to_string(),
success: false,
error_message: Some(resource),
context: Some(format!("peer_ip={peer_ip}")),
chain_hash: None,
});
}
let err = state.error_sanitizer.sanitize(GraphQLError::from_fraiseql_error(&e));
ErrorResponse::from_error(err)
})?;
let elapsed = start_time.elapsed();
#[allow(clippy::cast_possible_truncation)]
let elapsed_us = elapsed.as_micros() as u64;
metrics.queries_success.fetch_add(1, Ordering::Relaxed);
metrics.queries_duration_us.fetch_add(elapsed_us, Ordering::Relaxed);
metrics.db_queries_total.fetch_add(1, Ordering::Relaxed);
metrics.db_queries_duration_us.fetch_add(elapsed_us, Ordering::Relaxed);
metrics.operation_metrics.record(op_name, elapsed_us, false);
if let Some(cost) = estimated_cost {
metrics.queries_cost_total.fetch_add(cost, Ordering::Relaxed);
tracing::info!(
target: "fraiseql::cost_audit",
cost,
tenant = tenant_key.as_deref().unwrap_or(""),
operation = %op_name,
"operation cost"
);
}
#[cfg(feature = "federation")]
if fraiseql_core::federation::is_federation_query(&query) {
metrics.record_entity_resolution(elapsed_us, true);
}
debug!(
elapsed_ms = elapsed.as_millis(),
operation_name = ?request.operation_name,
"Query executed successfully"
);
#[allow(unused_mut)]
let mut response_json = result;
#[cfg(feature = "secrets")]
Box::pin(stages::decrypt_response_fields(&state, &mut response_json)).await?;
if let Some((key, body_hash)) = idempotency_key {
state
.idempotency_store
.store(
key,
body_hash,
crate::routes::idempotency::StoredResponse {
status: 200,
headers: Vec::new(),
body: Some(response_json.clone()),
},
)
.await;
}
Ok(GraphQLResponse {
body: response_json,
})
}
pub(super) fn tenant_dispatch_error(error: &FraiseQLError) -> GraphQLError {
match error {
FraiseQLError::ServiceUnavailable { retry_after, .. } => {
GraphQLError::service_unavailable(error.to_string(), *retry_after)
},
FraiseQLError::RateLimited { .. } => GraphQLError::rate_limited(error.to_string()),
FraiseQLError::CostExceeded {
retry_after_secs, ..
} => match retry_after_secs {
Some(secs) => GraphQLError::cost_budget_exhausted(error.to_string(), *secs),
None => GraphQLError::operation_cost_exceeded(error.to_string()),
},
_ => GraphQLError::new(error.to_string(), crate::error::ErrorCode::Forbidden),
}
}
mod incremental;
mod sse;
mod stages;
#[cfg(test)]
mod tests;