use axum::extract::{FromRequestParts, Query};
use axum::http::request::Parts;
use serde::Deserialize;
use toolkit_canonical_errors::CanonicalError;
use toolkit_odata::errors::OdataError;
use toolkit_odata::{CursorV1, Error as ODataError, ODataOrderBy, OrderKey, SortDir};
pub use toolkit_odata::ODataQuery;
#[derive(Deserialize, Default)]
pub struct ODataParams {
#[serde(rename = "$filter")]
pub filter: Option<String>,
#[serde(rename = "$orderby")]
pub orderby: Option<String>,
#[serde(rename = "$select")]
pub select: Option<String>,
#[serde(alias = "$top")]
pub limit: Option<u64>,
#[serde(alias = "$skiptoken")]
pub cursor: Option<String>,
}
pub const ACCEPTED_SYSTEM_QUERY_OPTIONS: [&str; 5] =
["$filter", "$orderby", "$select", "$top", "$skiptoken"];
const ODATA_SYSTEM_QUERY_OPTIONS: [&str; 12] = [
"$filter",
"$expand",
"$select",
"$orderby",
"$top",
"$skip",
"$count",
"$search",
"$format",
"$compute",
"$index",
"$schemaversion",
];
const UNSUPPORTED_QUERY_PARAM: &str = "UNSUPPORTED_QUERY_PARAM";
pub const MAX_FILTER_LEN: usize = 8 * 1024;
pub const MAX_NODES: usize = 2000;
pub const MAX_ORDERBY_LEN: usize = 1024;
pub const MAX_ORDER_FIELDS: usize = 10;
pub const MAX_SELECT_LEN: usize = 2048;
pub const MAX_SELECT_FIELDS: usize = 100;
fn select_invalid_arg(detail: impl Into<String>, reason: &'static str) -> CanonicalError {
OdataError::invalid_argument()
.with_field_violation("$select", detail, reason)
.create()
}
#[allow(clippy::result_large_err)]
pub fn parse_select(raw: &str) -> Result<Vec<String>, CanonicalError> {
let raw = raw.trim();
if raw.is_empty() {
return Err(select_invalid_arg(
"$select cannot be empty",
"INVALID_SELECT",
));
}
if raw.len() > MAX_SELECT_LEN {
return Err(select_invalid_arg("$select too long", "INVALID_SELECT"));
}
let fields: Vec<String> = raw
.split(',')
.map(|f| f.trim().to_lowercase())
.filter(|f| !f.is_empty())
.collect();
if fields.is_empty() {
return Err(select_invalid_arg(
"$select must contain at least one field",
"INVALID_SELECT",
));
}
if fields.len() > MAX_SELECT_FIELDS {
return Err(select_invalid_arg(
"$select contains too many fields",
"INVALID_SELECT",
));
}
let mut seen = std::collections::HashSet::new();
for field in &fields {
if !seen.insert(field.clone()) {
return Err(select_invalid_arg(
format!("duplicate field in $select: {field}"),
"INVALID_SELECT",
));
}
}
Ok(fields)
}
pub fn parse_orderby(raw: &str) -> Result<ODataOrderBy, toolkit_odata::Error> {
let raw = raw.trim();
if raw.is_empty() {
return Ok(ODataOrderBy::empty());
}
if raw.len() > MAX_ORDERBY_LEN {
return Err(toolkit_odata::Error::InvalidOrderByField(
"orderby too long".into(),
));
}
let mut keys = Vec::new();
for part in raw.split(',') {
let part = part.trim();
if part.is_empty() {
continue;
}
let tokens: Vec<&str> = part.split_whitespace().collect();
let (field, dir) = match tokens.as_slice() {
[field] | [field, "asc"] => (*field, SortDir::Asc),
[field, "desc"] => (*field, SortDir::Desc),
_ => {
return Err(toolkit_odata::Error::InvalidOrderByField(format!(
"invalid orderby clause: {part}"
)));
}
};
if field.is_empty() {
return Err(toolkit_odata::Error::InvalidOrderByField(
"empty field name in orderby".into(),
));
}
keys.push(OrderKey {
field: field.to_owned(),
dir,
});
}
if keys.len() > MAX_ORDER_FIELDS {
return Err(toolkit_odata::Error::InvalidOrderByField(
"too many order fields".into(),
));
}
Ok(ODataOrderBy(keys))
}
fn filter_invalid_arg(detail: impl Into<String>, reason: &'static str) -> CanonicalError {
OdataError::invalid_argument()
.with_field_violation("$filter", detail, reason)
.create()
}
fn query_params_invalid_arg(detail: impl Into<String>) -> CanonicalError {
OdataError::invalid_argument()
.with_field_violation("query", detail, "INVALID_QUERY_PARAMS")
.create()
}
fn unsupported_option_hint(option: &str) -> &'static str {
match option {
"$skip" => {
"this platform pages by cursor, not by offset: send the previous \
page's `next_cursor` as `$skiptoken` (alias `cursor`)"
}
"$count" => {
"a page carries no total: `PageInfo` is \
`{next_cursor, prev_cursor, limit}`"
}
"$format" => {
"responses are JSON: negotiate the media type with the `Accept` \
header"
}
_ => "not implemented by this platform",
}
}
fn unsupported_option_detail(option: &str) -> String {
let canonical = option.to_ascii_lowercase();
if ACCEPTED_SYSTEM_QUERY_OPTIONS.contains(&canonical.as_str()) {
return format!("`{option}` binds only as `{canonical}`");
}
if ODATA_SYSTEM_QUERY_OPTIONS.contains(&canonical.as_str()) {
return format!(
"unsupported OData system query option `{option}`; {}",
unsupported_option_hint(&canonical)
);
}
format!(
"unknown query option `{option}`; the `$` prefix is reserved for \
OData system query options"
)
}
#[allow(clippy::result_large_err)]
fn reject_unsupported_system_query_options(
pairs: &[(String, String)],
) -> Result<(), CanonicalError> {
let mut offending: Vec<&str> = Vec::new();
for key in pairs.iter().map(|(key, _)| key.as_str()) {
if key.starts_with('$')
&& !ACCEPTED_SYSTEM_QUERY_OPTIONS.contains(&key)
&& !offending.contains(&key)
{
offending.push(key);
}
}
let mut offenders = offending.into_iter();
let Some(first) = offenders.next() else {
return Ok(());
};
let mut violations = OdataError::invalid_argument().with_field_violation(
first,
unsupported_option_detail(first),
UNSUPPORTED_QUERY_PARAM,
);
for option in offenders {
violations = violations.with_field_violation(
option,
unsupported_option_detail(option),
UNSUPPORTED_QUERY_PARAM,
);
}
Err(violations.create())
}
pub async fn extract_odata_query<S>(
parts: &mut Parts,
state: &S,
) -> Result<ODataQuery, CanonicalError>
where
S: Send + Sync,
{
let Query(pairs) = Query::<Vec<(String, String)>>::from_request_parts(parts, state)
.await
.map_err(|e| query_params_invalid_arg(format!("Invalid query parameters: {e}")))?;
reject_unsupported_system_query_options(&pairs)?;
let Query(params) = Query::<ODataParams>::from_request_parts(parts, state)
.await
.map_err(|e| query_params_invalid_arg(format!("Invalid query parameters: {e}")))?;
let mut query = ODataQuery::new();
if let Some(raw_filter) = params.filter.as_ref() {
let raw = raw_filter.trim();
if !raw.is_empty() {
if raw.len() > MAX_FILTER_LEN {
return Err(filter_invalid_arg("Filter too long", "FILTER_TOO_LONG"));
}
let parsed = toolkit_odata::parse_filter_string(raw).map_err(|e| {
tracing::debug!(error = %e, filter_len = raw.len(), "OData filter parsing failed");
CanonicalError::from(e)
})?;
if parsed.node_count() > MAX_NODES {
tracing::debug!(
node_count = parsed.node_count(),
max_nodes = MAX_NODES,
"Filter complexity budget exceeded"
);
return Err(filter_invalid_arg(
"Filter too complex",
"FILTER_TOO_COMPLEX",
));
}
let filter_hash = toolkit_odata::pagination::short_filter_hash(Some(parsed.as_expr()));
let core_expr = parsed.into_expr();
query = query.with_filter(core_expr);
if let Some(hash) = filter_hash {
query = query.with_filter_hash(hash);
}
}
}
if params.cursor.is_some() && params.orderby.is_some() {
return Err(ODataError::OrderWithCursor.into());
}
if let Some(cursor_str) = params.cursor.as_ref() {
let cursor = CursorV1::decode(cursor_str).map_err(|_| ODataError::InvalidCursor)?;
query = query.with_cursor(cursor);
query = query.with_order(ODataOrderBy::empty());
} else if let Some(raw_orderby) = params.orderby.as_ref() {
let order = parse_orderby(raw_orderby).map_err(CanonicalError::from)?;
query = query.with_order(order);
}
if let Some(limit) = params.limit {
if limit == 0 {
return Err(ODataError::InvalidLimit.into());
}
query = query.with_limit(limit);
}
if let Some(raw_select) = params.select.as_ref() {
let fields = parse_select(raw_select)?;
query = query.with_select(fields);
}
Ok(query)
}
use std::ops::Deref;
#[derive(Debug, Clone)]
pub struct OData(pub ODataQuery);
impl OData {
#[inline]
pub fn into_inner(self) -> ODataQuery {
self.0
}
}
impl Deref for OData {
type Target = ODataQuery;
#[inline]
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl AsRef<ODataQuery> for OData {
#[inline]
fn as_ref(&self) -> &ODataQuery {
&self.0
}
}
impl From<OData> for ODataQuery {
#[inline]
fn from(x: OData) -> Self {
x.0
}
}
impl<S> FromRequestParts<S> for OData
where
S: Send + Sync,
{
type Rejection = CanonicalError;
#[allow(clippy::manual_async_fn)]
fn from_request_parts(
parts: &mut Parts,
state: &S,
) -> impl core::future::Future<Output = Result<Self, Self::Rejection>> + Send {
async move {
let query = extract_odata_query(parts, state).await?;
Ok(OData(query))
}
}
}
#[cfg(test)]
#[cfg_attr(coverage_nightly, coverage(off))]
#[path = "odata_tests.rs"]
mod odata_tests;