use std::collections::HashMap;
use axum::http::{HeaderMap, HeaderValue};
use fraiseql_core::{runtime::QueryMatch, security::SecurityContext};
use serde_json::json;
use super::{
RestHandler,
headers::{set_preference_applied, set_request_id},
prefer::PreferHeader,
response::{RestError, RestResponse},
routing::ResolvedGetQuery,
search::plan_search,
};
use crate::routes::rest::{
params::{ExtractedParams, PaginationParams, RestFieldSpec, RestParamExtractor},
resource::{HttpMethod, RouteSource},
response::helpers::{check_if_none_match, compute_etag},
};
pub fn refuse_unstreamable_request(
prefer: &PreferHeader,
params: &ExtractedParams,
) -> Result<(), RestError> {
let ExtractedParams {
path_params: _,
where_clause: _,
order_by: _,
pagination: _,
requested_pagination,
field_selection: _,
search_query: _,
embeddings,
embedding_filters,
embedding_pages,
embedding_counts,
} = params;
if prefer.count_exact || prefer.count_planned || prefer.count_estimated {
return Err(RestError::bad_request("count not available for export responses"));
}
if matches!(requested_pagination.offset, Some(offset) if offset > 0)
|| requested_pagination.after.is_some()
|| requested_pagination.last.is_some()
|| requested_pagination.before.is_some()
{
return Err(RestError::bad_request(
"pagination not available for export; use filters to narrow results",
));
}
if !embeddings.is_empty() {
let named = quoted_list(embeddings.iter().map(|spec| spec.relationship.as_str()));
return Err(RestError::bad_request(format!(
"embedded relationships are not available for export responses: {named}. An export \
is one statement over one snapshot; resolving an embed issues a sub-query per row. \
Request `Accept: application/json` to embed, or project the related data into the \
exported view."
)));
}
if !embedding_counts.is_empty() {
let named = quoted_list(embedding_counts.iter().map(|name| format!("{name}.count")));
return Err(RestError::bad_request(format!(
"embedded counts are not available for export responses: {named}. An export is one \
statement over one snapshot; a count issues a sub-query per row. Request \
`Accept: application/json` for counts."
)));
}
if !embedding_filters.is_empty() {
let named = quoted_list(embedding_filter_parameters(embedding_filters));
return Err(RestError::bad_request(format!(
"embedded-relationship filters are not available for export responses: {named}. A \
dotted parameter filters an embedded relationship, and an export carries no embed \
to filter. Narrow the exported rows themselves with `?field=value`, or request \
`Accept: application/json` to embed and filter."
)));
}
if !embedding_pages.is_empty() {
let named = quoted_list(embedding_pages.keys().map(|path| format!("{path}.limit")));
return Err(RestError::bad_request(format!(
"embedded-level pages are not available for export responses: {named}. A \
`rel.limit` pages an embedded relationship, and an export carries no embed to \
page. Bound the exported rows themselves with `?limit=`, or request \
`Accept: application/json` to embed."
)));
}
Ok(())
}
pub(super) fn refuse_unapplied_embedding_pages(
embeddings: &[super::super::params::EmbeddedSpec],
embedding_pages: &std::collections::BTreeMap<String, u32>,
lenient: bool,
) -> Result<(), RestError> {
if embedding_pages.is_empty() {
return Ok(());
}
let embedded = super::super::embedding::embedded_level_paths(embeddings);
let unapplied: Vec<String> = embedding_pages
.keys()
.filter(|path| !embedded.contains(path))
.map(|path| format!("{path}.limit"))
.collect();
if unapplied.is_empty() {
return Ok(());
}
let named = quoted_list(unapplied.iter());
if lenient {
tracing::debug!(
parameters = %named,
"ignoring embedded-level pages with no embed to apply them to \
(Prefer: handling=lenient)"
);
return Ok(());
}
Err(RestError::bad_request(format!(
"pages were sent for levels this request did not embed: {named}. `rel.limit` pages \
the rows of an embedded relationship — embed it in `?select=` (a nested level by its \
dotted path, `?orders.items.limit=` for `orders(items(...))`), or drop the parameter. \
`Prefer: handling=lenient` ignores such a page instead."
)))
}
pub(super) fn refuse_unapplied_embedding_filters(
embeddings: &[super::super::params::EmbeddedSpec],
embedding_counts: &[String],
embedding_filters: &HashMap<String, serde_json::Value>,
lenient: bool,
) -> Result<(), RestError> {
if embedding_filters.is_empty() {
return Ok(());
}
let applied: std::collections::HashSet<&str> = embeddings
.iter()
.map(|spec| spec.relationship.as_str())
.chain(embedding_counts.iter().map(String::as_str))
.collect();
let unapplied: HashMap<String, serde_json::Value> = embedding_filters
.iter()
.filter(|(relationship, _)| !applied.contains(relationship.as_str()))
.map(|(k, v)| (k.clone(), v.clone()))
.collect();
if unapplied.is_empty() {
return Ok(());
}
let named = quoted_list(embedding_filter_parameters(&unapplied));
if lenient {
tracing::debug!(
parameters = %named,
"ignoring embedded-relationship filters with no embed to apply them to \
(Prefer: handling=lenient)"
);
return Ok(());
}
let mut relationships: Vec<&str> = unapplied.keys().map(String::as_str).collect();
relationships.sort_unstable();
let first = relationships.first().copied().unwrap_or("rel");
Err(RestError::bad_request(format!(
"filters were sent for relationships this request did not embed: {named}. A dotted \
parameter narrows an embedded relationship — add `{first}(...)` to `?select=` to embed \
and filter it, or `{first}.count` to count the matching rows, or drop the filter. \
`Prefer: handling=lenient` ignores such a filter instead."
)))
}
fn embedding_filter_parameters(filters: &HashMap<String, serde_json::Value>) -> Vec<String> {
let mut named: Vec<String> = filters
.iter()
.flat_map(|(relationship, fields)| match fields.as_object() {
Some(obj) if !obj.is_empty() => {
obj.keys().map(|field| format!("{relationship}.{field}")).collect::<Vec<_>>()
},
_ => vec![relationship.clone()],
})
.collect();
named.sort();
named
}
fn quoted_list<S: AsRef<str>>(names: impl IntoIterator<Item = S>) -> String {
names
.into_iter()
.map(|n| format!("`{}`", n.as_ref()))
.collect::<Vec<_>>()
.join(", ")
}
impl RestHandler<'_> {
pub fn resolve_streaming_get_query(
&self,
relative_path: &str,
query_pairs: &[(&str, &str)],
headers: &http::HeaderMap,
security_context: Option<&SecurityContext>,
) -> Result<ResolvedGetQuery, RestError> {
let resolved =
self.resolve_get_query(relative_path, query_pairs, headers, security_context)?;
let streamable = self
.schema
.find_query(&resolved.query_name)
.is_some_and(|query_def| query_def.rest_stream);
if !streamable {
return Err(RestError::not_acceptable(format!(
"`{}` is not exported as a stream; set `rest_stream = true` on the query to \
offer NDJSON, CSV and XLSX on this route",
resolved.query_name
)));
}
refuse_unstreamable_request(&PreferHeader::from_headers(headers), &resolved.params)?;
Ok(resolved)
}
pub fn resolve_get_query(
&self,
relative_path: &str,
query_pairs: &[(&str, &str)],
headers: &http::HeaderMap,
security_context: Option<&SecurityContext>,
) -> Result<ResolvedGetQuery, RestError> {
let resolved = self
.route_table
.resolve(relative_path, HttpMethod::Get)
.ok_or_else(|| RestError::not_found("Route not found"))?;
let query_name = match &resolved.route.source {
RouteSource::Query { name } => name.as_str(),
RouteSource::Mutation { .. } => {
return Err(RestError::internal("GET route backed by mutation"));
},
};
let query_def = self
.schema
.find_query(query_name)
.ok_or_else(|| RestError::not_found(format!("Query not found: {query_name}")))?;
let type_def = self.schema.find_type(&query_def.return_type);
let lenient = PreferHeader::from_headers(headers).handling
== Some(super::prefer::HandlingPreference::Lenient);
let extractor = RestParamExtractor::new(self.config, query_def, type_def)
.with_lenient_handling(lenient);
let path_pairs: Vec<(&str, &str)> =
resolved.path_params.iter().map(|(k, v)| (k.as_str(), v.as_str())).collect();
let params = extractor.extract(&path_pairs, query_pairs)?;
let field_names = match ¶ms.field_selection {
RestFieldSpec::All => type_def
.map(|t| t.fields.iter().map(|f| f.name.to_string()).collect())
.unwrap_or_default(),
RestFieldSpec::Fields(fields) => fields.clone(),
};
let mut arguments: HashMap<String, serde_json::Value> = HashMap::new();
for (key, value) in ¶ms.path_params {
arguments.insert(key.clone(), value.clone());
}
let search = match params.search_query.as_deref() {
None => None,
Some(query) => Some(
plan_search(query, type_def, self.schema.security.as_ref(), security_context)
.ok_or_else(|| RestError {
status: http::StatusCode::FORBIDDEN,
code: "FORBIDDEN",
message: "`?search=` has no field this request may read to search"
.to_string(),
details: None,
})?,
),
};
let fts_where = search.as_ref().map(|plan| plan.where_clause.clone());
match (¶ms.where_clause, &fts_where) {
(Some(regular), Some(fts)) => {
arguments.insert("where".to_string(), json!({ "_and": [regular, fts] }));
},
(Some(regular), None) => {
arguments.insert("where".to_string(), regular.clone());
},
(None, Some(fts)) => {
arguments.insert("where".to_string(), fts.clone());
},
(None, None) => {},
}
let relevance = if let Some(ref order_by) = params.order_by {
arguments.insert("orderBy".to_string(), order_by.clone());
None
} else {
search.map(|plan| plan.relevance)
};
if let PaginationParams::Offset { limit, offset } = ¶ms.pagination {
arguments.insert("limit".to_string(), json!(limit));
if *offset > 0 {
arguments.insert("offset".to_string(), json!(offset));
}
}
let mut variables = serde_json::Map::new();
for (k, v) in &arguments {
variables.insert(k.clone(), v.clone());
}
if let PaginationParams::Cursor {
first,
after,
last,
before,
} = ¶ms.pagination
{
if let Some(f) = first {
variables.insert("first".to_string(), json!(f));
}
if let Some(ref a) = after {
variables.insert("after".to_string(), json!(a));
}
if let Some(l) = last {
variables.insert("last".to_string(), json!(l));
}
if let Some(ref b) = before {
variables.insert("before".to_string(), json!(b));
}
}
let variables_json = serde_json::Value::Object(variables);
let mut query_match =
QueryMatch::from_operation(query_def.clone(), field_names, arguments, type_def)?;
if let Some(relevance) = relevance {
query_match = query_match.with_search_relevance(relevance);
}
Ok(ResolvedGetQuery {
query_name: query_name.to_string(),
query_match,
variables: variables_json,
params,
})
}
pub async fn handle_get(
&self,
relative_path: &str,
query_pairs: &[(&str, &str)],
headers: &HeaderMap,
security_context: Option<&SecurityContext>,
) -> Result<RestResponse, RestError> {
let resolved_query =
self.resolve_get_query(relative_path, query_pairs, headers, security_context)?;
let query_match = &resolved_query.query_match;
let variables_json = &resolved_query.variables;
let params = &resolved_query.params;
let prefer = PreferHeader::from_headers(headers);
let vars_ref = if variables_json.as_object().is_none_or(|m| m.is_empty()) {
None
} else {
Some(variables_json)
};
refuse_unapplied_embedding_filters(
¶ms.embeddings,
¶ms.embedding_counts,
¶ms.embedding_filters,
prefer.handling == Some(super::prefer::HandlingPreference::Lenient),
)?;
refuse_unapplied_embedding_pages(
¶ms.embeddings,
¶ms.embedding_pages,
prefer.handling == Some(super::prefer::HandlingPreference::Lenient),
)?;
let request_budget = self.executor.request_budget();
let has_embeddings = !params.embeddings.is_empty() || !params.embedding_counts.is_empty();
let (embeds, counts) = super::super::embedding::selections(
¶ms.embeddings,
¶ms.embedding_counts,
¶ms.embedding_filters,
¶ms.embedding_pages,
super::super::embedding::default_embed_page(self.config),
);
let read = async {
if has_embeddings {
self.executor
.execute_query_composed(
query_match,
&embeds,
&counts,
vars_ref,
security_context,
Some(&request_budget),
)
.await
} else {
self.executor
.execute_query_direct(
query_match,
vars_ref,
security_context,
Some(&request_budget),
)
.await
}
};
let (result, total, count_applied) = if prefer.count_preference().is_some() {
let (r, c) = tokio::join!(
read,
self.executor.count_rows(query_match, vars_ref, security_context),
);
(r?, Some(c?), Some("count=exact"))
} else {
(read.await?, None, None)
};
let mut response_headers = HeaderMap::new();
set_request_id(headers, &mut response_headers);
let mut applied: Vec<&str> = Vec::new();
if let Some(count_pref) = count_applied {
applied.push(count_pref);
}
if prefer.handling == Some(super::prefer::HandlingPreference::Lenient) {
applied.push("handling=lenient");
}
if !applied.is_empty() {
set_preference_applied(&mut response_headers, &applied);
}
if (prefer.count_planned || prefer.count_estimated) && count_applied == Some("count=exact")
{
response_headers
.insert("x-preference-fallback", HeaderValue::from_static("count=exact"));
}
super::super::cache_control::apply_cache_headers(
&mut response_headers,
&super::super::cache_control::CacheContext {
is_get: true,
authenticated: security_context.is_some(),
query_ttl: query_match.query_def.cache_ttl_seconds,
default_ttl: self.config.default_cache_ttl,
cdn_max_age: self.config.cdn_max_age,
},
);
let body = build_query_response(&result, total, ¶ms.pagination)?;
if self.config.etag {
let serialized = serde_json::to_vec(&body).map_err(|e| {
RestError::internal(format!("Failed to serialize response for ETag: {e}"))
})?;
let etag = compute_etag(&serialized);
if check_if_none_match(headers, &etag).unwrap_or(false) {
let mut not_modified = response_headers.clone();
not_modified.insert(
"etag",
HeaderValue::from_str(&etag).map_err(|e| {
RestError::internal(format!("Computed ETag is not a valid header: {e}"))
})?,
);
return Ok(RestResponse {
status: axum::http::StatusCode::NOT_MODIFIED,
headers: not_modified,
body: None,
});
}
response_headers.insert(
"etag",
HeaderValue::from_str(&etag).map_err(|e| {
RestError::internal(format!("Computed ETag is not a valid header: {e}"))
})?,
);
}
Ok(RestResponse {
status: axum::http::StatusCode::OK,
headers: response_headers,
body: Some(body),
})
}
}
pub(super) fn build_query_response(
result: &serde_json::Value,
total: Option<u64>,
pagination: &PaginationParams,
) -> Result<serde_json::Value, RestError> {
let data = if let Some(data_obj) = result.get("data") {
if let serde_json::Value::Object(map) = data_obj {
map.values().next().cloned().unwrap_or(serde_json::Value::Null)
} else {
data_obj.clone()
}
} else {
result.clone()
};
let mut response = json!({ "data": data });
match pagination {
PaginationParams::Offset { limit, offset } => {
let mut meta = json!({
"limit": limit,
"offset": offset,
});
if let Some(total) = total {
meta["total"] = json!(total);
}
response["meta"] = meta;
},
PaginationParams::Cursor {
first,
after,
last,
before,
} => {
let mut meta = serde_json::Map::new();
if let Some(page_info) = extract_relay_page_info(&data) {
if let Some(has_next) = page_info.get("hasNextPage") {
meta.insert("hasNextPage".to_string(), has_next.clone());
}
if let Some(has_prev) = page_info.get("hasPreviousPage") {
meta.insert("hasPreviousPage".to_string(), has_prev.clone());
}
}
if let Some(f) = first {
meta.insert("first".to_string(), json!(f));
}
if let Some(ref a) = after {
meta.insert("after".to_string(), json!(a));
}
if let Some(l) = last {
meta.insert("last".to_string(), json!(l));
}
if let Some(ref b) = before {
meta.insert("before".to_string(), json!(b));
}
if let Some(total) = total {
meta.insert("total".to_string(), json!(total));
}
response["meta"] = serde_json::Value::Object(meta);
},
PaginationParams::None => {
},
}
Ok(response)
}
pub(super) fn extract_relay_page_info(data: &serde_json::Value) -> Option<&serde_json::Value> {
data.get("pageInfo")
}